feat: 增加 MCP 研究模板保存工具

This commit is contained in:
yuxuanhui
2026-09-11 23:05:01 +08:00
parent 03a66546f3
commit f29063c9a2
13 changed files with 430 additions and 13 deletions
+1 -1
View File
@@ -11,7 +11,7 @@ from sqlalchemy import select
from ..models import Account, Admin, MCPToken, now
from ..security import token_hash
SCOPES = frozenset({"research:read", "research:refresh", "backtests:execute", "backtests:control"})
SCOPES = frozenset({"research:read", "research:refresh", "research:write", "backtests:execute", "backtests:control"})
@dataclass(frozen=True)
+3 -2
View File
@@ -21,6 +21,7 @@ from ..research_access.service import ResearchAccess, ResearchError
# Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
TOOLS = {
"create_research_template": (c.CreateTemplate, "create_template", "research:write", "将调用方大模型研究后自行总结的参数化模板保存到模板工坊,供用户后续批量回测。先用 get_backtest_results 阅读实际指标和检查,选择 1–20 个已完成采集的 source_item_ids,并说明 hypothesis;不要把 completed 当作检查通过。template 使用 {name} 占位符及逐一对应的 variables,字段变量须声明 MATRIX/VECTOR/GROUP,VECTOR 聚合须明确写入表达式。提供唯一名称和 idempotency_key,可附 reference。返回模板 ID、版本和理论组合数;仅核验结构及来源,不验证所有参数组合,不再次调用模型、不执行回测、不覆盖已有模板。"),
"get_submission_check": (c.SelfCorrelationReference, "submission_check_context", "research:read", "读取已导入 Alpha 的表达式、Description、snapshot 和缓存检查结果;不发起检查。先核对或生成三段 Description,再调用 check_submission。"),
"check_submission": (c.SubmissionCheck, "check_submission", "research:refresh", "对单个待提交 Alpha 写回已获用户授权的 Description 并调用平台 GET /check,返回 job_id。须先用 get_submission_check 获取 snapshot;保留本地自相关门槛和冲突保护。通过 get_refresh_job 查进度、get_submission_check 读结果。无论检查结果如何,都不会调用 /submit 或正式提交 Alpha。"),
"get_worldquant_connection": (c.ConnectionReference, "connection", "research:read", "读取 WorldQuant 连接状态及可选认证 job_id 的进度,不发起认证;人工验证在网页完成。"),
@@ -69,7 +70,7 @@ class MCPResearchServer:
return types.ListToolsResult(tools=[types.Tool(name=name, description=description,
inputSchema=schema.model_json_schema(), annotations=types.ToolAnnotations(
readOnlyHint=scope == "research:read", destructiveHint=method == "control",
idempotentHint=method in {"submit", "control"} or scope == "research:read",
idempotentHint=method in {"submit", "control", "create_template"} or scope == "research:read",
openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation", "check_submission", "authenticate"}))
for name, (schema, method, scope, description) in TOOLS.items()
if scope in principal.scopes and "research:read" in principal.scopes])
@@ -123,7 +124,7 @@ class MCPResearchServer:
data = {"error": ResearchError(code, "研究操作失败;可使用原幂等键重试或查询历史", retryable=True).data}
db.add(MCPAudit(id=str(uuid4()), token_id=principal.token_id, tool=name,
request_id=fingerprint({"request_id": request_id}), input_digest=digest,
business_id=data.get("backtest_run_id", data.get("job_id")),
business_id=data.get("backtest_run_id", data.get("job_id", data.get("template_id"))),
result_code=code, elapsed_ms=int((time.monotonic()-started)*1000)))
if not error:
if access.wake == "backtests":
+2 -2
View File
@@ -8,7 +8,7 @@ from sqlalchemy import func, select
from ..models import Account, MCPToken, now
from ..security import require_auth
from .auth import create_token
from .auth import SCOPES, create_token
router = APIRouter(prefix="/api/v1/mcp-tokens", tags=["mcp-tokens"], dependencies=[Depends(require_auth)])
@@ -17,7 +17,7 @@ class TokenInput(BaseModel):
model_config = ConfigDict(extra="forbid")
name: str = Field(min_length=1, max_length=100)
days: int = Field(default=90, ge=1, le=365, strict=True)
scopes: list[str] = Field(default_factory=lambda: ["research:read"], max_length=4)
scopes: list[str] = Field(default_factory=lambda: ["research:read"], max_length=len(SCOPES))
def token_output(row, account):
+21
View File
@@ -7,6 +7,7 @@ from pydantic import Field, model_validator
from ..backtests.contracts import Candidate, SimulationSettings
from ..catalog.contracts import CatalogFilters, Scope
from ..research.workspace_contracts import TemplateSpec
from ..schemas import Contract
Identifier = Annotated[str, Field(min_length=1, max_length=100)]
@@ -74,6 +75,26 @@ class Control(Contract):
idempotency_key: Identifier
class CreateTemplate(Contract):
template: TemplateSpec
hypothesis: str = Field(min_length=1, max_length=10000)
source_item_ids: list[RunId] = Field(min_length=1, max_length=20)
reference: str | None = Field(default=None, max_length=200)
idempotency_key: Identifier
@model_validator(mode="after")
def research_template(self):
self.template.name = self.template.name.strip()
self.hypothesis = self.hypothesis.strip()
if not self.template.name or not self.hypothesis or not self.template.expression.strip():
raise ValueError("模板名称、表达式和研究假设不能为空")
if self.template.category != "template":
raise ValueError("此工具仅创建完整模板,不保存表达式片段")
if len(set(self.source_item_ids)) != len(self.source_item_ids):
raise ValueError("来源候选 ID 不能重复")
return self
class CatalogSearch(Contract):
filters: CatalogFilters
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
+25 -5
View File
@@ -53,6 +53,12 @@ class ResearchAccess:
"confirmation": "调用者须已获本批执行授权;直接提交后返回稳定运行 ID",
"duplicate_policies": ["reject", "rerun"], "permissions": sorted(self.principal.scopes),
"metadata_only": True, "actual_platform_allowance": None,
"templates": {
"create_with": "create_research_template", "required_scope": "research:write",
"authored_by": "caller", "max_source_items": 20,
"source_items_with": "get_backtest_results", "starts_backtests": False,
"web_url": f"{self.public_origin}/#templates",
},
"submission_check": {
"check_with": "check_submission", "read_with": "get_submission_check",
"job_with": "get_refresh_job", "max_targets": 1,
@@ -220,6 +226,18 @@ class ResearchAccess:
async def history(self, args):
return await self.evidence.history(args)
async def create_template(self, args):
"""Persist the caller's template and evidence without model or queue execution."""
from .templates import create_template
previous, digest = await self.previous("create_research_template", args)
if previous:
return previous.response
result = await create_template(self.db, args, self.principal)
result["web_url"] = f"{self.public_origin}/#templates"
return await self.remember("create_research_template", args, digest, result,
business_id=result["template_id"])
async def previous(self, operation, args):
# PostgreSQL row lock is shared with HTTP start and catalog/job creation.
account = await self.db.scalar(select(Account).where(Account.id == self.principal.account_id).with_for_update())
@@ -233,13 +251,13 @@ class ResearchAccess:
raise ResearchError("IDEMPOTENCY_CONFLICT", "幂等键已用于不同内容")
return row, digest
async def remember(self, operation, args, digest, result):
async def remember(self, operation, args, digest, result, *, business_id, wake=None):
result["_meta"] = {"schema_version": 1, "observed_at": now().isoformat(), "source": "system"}
self.db.add(ResearchRequest(id=str(uuid4()), account_id=self.principal.account_id,
operation=operation, idempotency_key=args.idempotency_key, digest=digest,
business_id=result["backtest_run_id"], response=encode_snapshot(result)))
business_id=business_id, response=encode_snapshot(result)))
await self.db.flush()
self.wake = "backtests"
self.wake = wake
return result
async def validate_settings(self, candidates):
@@ -286,7 +304,8 @@ class ResearchAccess:
result = {**result, "input_digest": digest, "batch_count": preview["batch_count"],
"duplicates": {"within_batch": within, "historical_matches": history["total"]},
"validation": validation, "web_url": self.run_url(result["backtest_run_id"])}
return await self.remember("submit_backtests", args, digest, result)
return await self.remember("submit_backtests", args, digest, result,
business_id=result["backtest_run_id"], wake="backtests")
async def run(self, args):
result = await self.backtests.run(args.run_id)
@@ -325,4 +344,5 @@ class ResearchAccess:
"indefinite_account_block_cleared": args.action == "resume" and bool(before["scheduler"]["blocked_reason"])
and before["scheduler"]["blocked_until"] is None,
"note": "暂停/停止仅阻止后续提交;已提交模拟继续采集。recover 不重新提交。"}
return await self.remember("control_backtest", args, digest, result)
return await self.remember("control_backtest", args, digest, result,
business_id=result["backtest_run_id"], wake="backtests")
+60
View File
@@ -0,0 +1,60 @@
"""Persist caller-authored templates with frozen local research evidence."""
import math
from sqlalchemy import select
from ..models import BacktestItem, BacktestResult, ResearchAsset
from ..research.assets import Assets
from ..research.workspace_contracts import AssetWrite
from .queries import item_summary
async def create_template(db, args, principal):
"""Save a new asset inside the caller's account-locked, idempotent transaction.
Args contain the external model's TemplateSpec and local source item IDs.
Return the versioned asset and theoretical combination count. Raise
ResearchError for name conflicts or missing/incomplete research evidence.
Stored evidence proves provenance, not profitability or platform eligibility;
schema validation does not validate every expanded FASTEXPR combination.
"""
from .service import ResearchError
existing = await db.scalar(select(ResearchAsset).where(
ResearchAsset.kind == "template", ResearchAsset.name == args.template.name,
).order_by(ResearchAsset.id).limit(1))
if existing:
raise ResearchError("TEMPLATE_NAME_CONFLICT", "模板名称已存在,请使用新名称;此工具不覆盖已有模板",
affected_items=[{"template_id": existing.id, "version": existing.version}])
rows = (await db.execute(select(BacktestItem, BacktestResult).outerjoin(
BacktestResult, BacktestResult.item_id == BacktestItem.id,
).where(BacktestItem.id.in_(args.source_item_ids)))).all()
found = {item.id: (item, result) for item, result in rows}
missing = [item_id for item_id in args.source_item_ids if item_id not in found]
if missing:
raise ResearchError("NOT_FOUND", "部分来源回测候选不存在", affected_items=missing)
incomplete = [item.id for item, result in rows if (
item.platform_status != "completed" or item.collection_status != "complete"
or item.persistence_status != "saved" or not result or not result.complete
)]
if incomplete:
raise ResearchError("SOURCE_NOT_READY", "来源候选须完成回测、结果采集和持久化,请先读取 get_backtest_results",
affected_items=sorted(incomplete))
provenance = {
"source": {"kind": "mcp", "reference": args.reference},
"hypothesis": args.hypothesis,
"mcp_token_id": principal.token_id,
"admin_id": principal.admin_id,
"source_items": [item_summary(*found[item_id]) for item_id in args.source_item_ids],
}
asset = await Assets(db).save(AssetWrite(
kind="template", content=args.template.model_dump(mode="json"),
), provenance=provenance)
return {
**asset, "template_id": asset["id"],
"combination_count": str(math.prod(len(v.values) for v in args.template.variables.values())),
"validation": {"structure": "valid", "source_evidence": "recorded",
"expanded_candidates": "not_validated", "platform_semantics": "unknown"},
"next_step": "在模板工坊选择固定输入及模拟设置,展开并核验候选,再确认批量回测。",
}