153 lines
17 KiB
Python
153 lines
17 KiB
Python
"""MCP transport over shared research operations, with minimal durable audit evidence."""
|
||
|
||
import asyncio
|
||
import json
|
||
import time
|
||
from uuid import uuid4
|
||
|
||
import anyio
|
||
from fastapi import HTTPException
|
||
from mcp import types
|
||
from mcp.server.lowlevel import Server
|
||
from mcp.server.transport_security import TransportSecuritySettings
|
||
from pydantic import ValidationError
|
||
|
||
from ..alphas import sanitize
|
||
from ..backtests.contracts import fingerprint
|
||
from ..models import MCPAudit, now
|
||
from ..research.serialization import encode_snapshot
|
||
from ..research_access import contracts as c
|
||
from ..research_access.service import ResearchAccess, ResearchError
|
||
from ..superalpha import contracts as sc
|
||
|
||
# Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
|
||
TOOLS = {
|
||
"search_superalpha_plans": (sc.PlanSearch, "super_plans", "research:read", "分页查找 Super Alpha 研究方案。"),
|
||
"get_superalpha_plan": (sc.PlanReference, "super_plan", "research:read", "读取指定方案版本或固定构造记录;不发起回测。"),
|
||
"save_superalpha_plan": (sc.PlanSave, "save_super_plan", "research:write", "保存调用方构造的 Selection/Combo 参数方案;更新须携带版本,支持幂等。不调用模型或回测。"),
|
||
"preview_superalpha_selection": (sc.SelectionPreview, "preview_super_selection", "research:refresh", "主动预览展开后的 Selection;异步返回 job_id,用 get_refresh_job 查进度、get_superalpha_selection 查组件。预览不是实际回测组件。"),
|
||
"get_superalpha_selection": (sc.SelectionReference, "super_selection", "research:read", "分页读取组件预览及完整性、时间、警告;缺失不自动刷新。"),
|
||
"build_superalpha_candidates": (sc.BuildCandidates, "build_super_candidates", "research:write", "按方案版本或内联方案进行全量展开/固定种子采样;保存固定候选及来源,不执行回测。将 candidates 与 submit_source 交给 submit_backtests;超过100项按分页读取固定记录。"),
|
||
"search_superalphas": (sc.SuperAlphaSearch, "super_alphas", "research:read", "分页查询本地已导入的 SUPER 成果,固定 SUPER 范围;不自动同步。"),
|
||
"get_superalpha": (sc.AlphaReference, "super_alpha", "research:read", "读取已导入 SUPER 的 Selection/Combo、指标、组件证据、Description 和研究来源。"),
|
||
"get_pyramid_distribution": (c.PyramidQuery, "pyramid_distribution", "research:read", "实时读取指定 region(如 USA、GLB)和 delay(0/1)的个人 Pyramid Alpha 分布;必传 current_date(YYYY-MM-DD),自动按自然年四季度取完整起止日(如2026-09-13对应2026-07-01至2026-09-30),传给平台 startDate/endDate,不使用默认周期。按用户约定 alphaCount>=3 为 lit(已点亮),1–2 为 in_progress,0 为 unlit;每项含 category、alpha_count、距3条的 remaining。复用平台认证,未连接时先调用 authenticate_worldquant;缺失数据不当作0。不回测、不提交。"),
|
||
"search_data_preparations": (c.PreparationSearch, "preparations", "research:read", "分页查询数据准备集合,返回固定范围、字段数及版本。研究可使用多个集合,各集合范围独立。"),
|
||
"get_data_preparation": (c.PreparationRead, "preparation", "research:read", "按集合 ID 与版本分页预览字段、类型、描述和数据集归属。提交回测时携带 preparation_refs,由服务端核对版本并固定独立输入快照;空集合不能用于研究。"),
|
||
"search_research_templates": (c.TemplateSearch, "templates", "research:read", "分页搜索模板工坊的模板与最新版本,不执行研究。"),
|
||
"get_research_template": (c.TemplateRead, "template", "research:read", "读取模板内容、字段定义和来源;新增版本前核对最新版本。"),
|
||
"create_research_template_version": (c.CreateTemplateVersion, "create_template_version", "research:write", "为已有模板新增不可变版本。先读取模板,携带 template_id、expected_version、完整 template、hypothesis 和 idempotency_key。source_item_ids 可选,提供时须为真实、已完成采集的回测项。版本冲突须重新读取;不调用模型、不回测。"),
|
||
"create_research_template": (c.CreateTemplate, "create_template", "research:write", "保存调用方编写的模板,不要求已有回测结果。提供 template、hypothesis、唯一名称和 idempotency_key;source_item_ids 可选,提供时须为真实、已完成采集的回测项。template 使用 {name} 占位符及对应 variables,字段定义只需类型和描述;空字段 values 由展开时的数据准备绑定,其他空参数须补充 values 或直接写入表达式。仅保存模板,不调用模型、不展开、不执行回测。"),
|
||
"expand_research_template": (c.TemplateExpansion, "expand_template", "research:write", "按 template_id/version、preparation_refs 和完整 settings 生成固定候选集合,支持全组合或固定 seed 随机采样。仅检查语法和数据准备/回测参数组合一致性,无逐行校验状态,不回测。返回首25项和 experiment_id,更多候选用 get_template_candidates 分页读取;用户授权后按明确 candidate_ids 调用 start_template_backtest。重试复用 idempotency_key。"),
|
||
"get_template_candidates": (c.TemplateCandidates, "template_candidates", "research:read", "分页读取固定模板候选集合的表达式、参数、候选 ID、模板版本、数据准备及关联回测;每页最多100条,total 不是当前页数量。"),
|
||
"start_template_backtest": (c.SubmitTemplateBacktest, "start_template_backtest", "backtests:execute", "对用户已授权的模板候选集合执行批量回测。提供 experiment_id、明确的 candidate_ids 和 idempotency_key;服务端使用保存的表达式和参数并保留来源,不接受重写输入。不需要另建预览,不重新校验类型或平台可用性。相同请求重试返回同一运行;新幂等键会创建新的回测,包括重复候选。立即返回运行 ID,结果另行查询。"),
|
||
"get_submission_check": (c.SelfCorrelationReference, "submission_check_context", "research:read", "读取已导入 Alpha 的表达式、Description、snapshot 和缓存检查结果;check_summary 分离 Alpha 检查和 REGULAR_SUBMISSION 提交限制,原始 checks 保留;限制不代表当前额度。不发起检查。先核对或生成三段 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 的进度,不发起认证;人工验证在网页完成。"),
|
||
"authenticate_worldquant": (c.Authentication, "authenticate", "research:refresh", "使用服务端已保存凭据连接或重新认证 WorldQuant,返回 job_id;action=connect(默认)或人工验证后 verify。用 get_worldquant_connection 查询,不接收密码,不修改账户配置。"),
|
||
"get_research_capabilities": (c.Empty, "capabilities", "research:read", "读取直接研究能力、完整设置 schema 和调度阻塞,不代表平台剩余额度。"),
|
||
"search_catalog": (c.CatalogSearch, "catalog", "research:read", "分页查询指定范围的数据集或字段元数据;无缓存不等于无数据,不隐式刷新。"),
|
||
"get_research_metadata": (c.Metadata, "metadata", "research:read", "读取范围、设置快照、算子定义或字段可用性;未知不认定通过。"),
|
||
"refresh_research_data": (c.Refresh, "refresh", "research:refresh", "显式刷新目录、算子、设置、字段可用性或 PnL;不会创建模拟。任务返回 job_id。"),
|
||
"get_refresh_job": (c.JobReference, "refresh_job", "research:read", "查询研究刷新、本地自相关或平台检查任务的状态、进度与分页错误。"),
|
||
"check_self_correlation": (c.SelfCorrelationCheck, "check_self_correlation", "research:refresh", "对 1–100 个已导入 Alpha 发起本地自相关检查,返回 job_id。与本地同地区已提交 Alpha 比较,排除自身;优先用缓存,缺失 PnL 自动补取。需先同步已提交 Alpha;不调用平台提交检查。用 get_refresh_job 查进度、get_self_correlation 读结果。"),
|
||
"get_self_correlation": (c.SelfCorrelationReference, "self_correlation", "research:read", "读取指定 Alpha 最新的本地自相关缓存,包括最大相关系数、样本覆盖与 stale 状态;无缓存不自动检查,结果不等同于平台提交资格。"),
|
||
"search_backtests": (c.History, "history", "research:read", "分页查历史候选与固定设置;candidates 按完整输入精确匹配,不推断数学等价。"),
|
||
"submit_backtests": (c.Submit, "submit", "backtests:execute", "执行用户已授权的 REGULAR/SUPER 固定批次,自动留痕并立即返回运行 ID;SUPER 使用 selection/combo 和专属设置,逐条模拟。可携带 preparation_refs 选择集合,版本变化须重新读取;每项必须完整设置;重复默认拒绝,rerun 明确重跑。不需要研究资产。"),
|
||
"get_backtest": (c.RunReference, "run", "research:read", "读取真实运行进度、提交数量和可选增量事件;受理不等于成功。"),
|
||
"get_backtest_results": (c.Results, "results", "research:read", "分页读取固定快照指标、Alpha 非通过检查及三层状态;REGULAR_SUBMISSION 单列 submission_limits,不计入 Alpha 失败统计。缺失指标不补零。"),
|
||
"get_backtest_artifact": (c.Artifact, "artifact", "research:read", "分页读取候选脱敏快照的顶层键值或独立采集的 PnL;缺缓存不自动刷新。"),
|
||
"control_backtest": (c.Control, "control", "backtests:control", "对已授权运行暂停、继续、停止或恢复采集;不远程取消、不重提未知模拟。需要版本和幂等键。"),
|
||
}
|
||
|
||
|
||
def tool_result(data, error=False):
|
||
data = encode_snapshot(data)
|
||
return types.CallToolResult(content=[types.TextContent(type="text", text=json.dumps(data, ensure_ascii=False))],
|
||
structuredContent=data, isError=error)
|
||
|
||
|
||
class MCPResearchServer:
|
||
def __init__(self, sessions, runner, settings):
|
||
self.sessions, self.runner, self.settings = sessions, runner, settings
|
||
# The existing deployment has one owner; this also gives SQLite test transactions a fair queue.
|
||
self.mutation_lock = asyncio.Lock()
|
||
self.server = Server("wq-alpha-research", version="1.0.0", on_list_tools=self.list_tools,
|
||
on_call_tool=self.call_tool,
|
||
instructions="自由探索,直接固定候选回测,无需先建研究资产。使用已有模板时,expand_research_template 生成候选,get_template_candidates 分页核对,获得用户授权后 start_template_backtest 执行;无需额外预览或结果评估步骤。工具不安排定时研究;结果按运行 ID 查询。")
|
||
from urllib.parse import urlsplit
|
||
|
||
host = urlsplit(settings.public_origin).netloc
|
||
self.app = self.server.streamable_http_app(
|
||
streamable_http_path="/", stateless_http=True, json_response=True,
|
||
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=True,
|
||
allowed_hosts=[host], allowed_origins=[settings.public_origin.rstrip("/")]),
|
||
)
|
||
|
||
async def list_tools(self, ctx, params):
|
||
principal = ctx.request.state.mcp_principal
|
||
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", "create_template", "create_template_version", "save_super_plan", "build_super_candidates", "expand_template", "start_template_backtest"} or scope == "research:read",
|
||
openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation", "check_submission", "authenticate", "pyramid_distribution", "preview_super_selection", "start_template_backtest"}))
|
||
for name, (schema, method, scope, description) in TOOLS.items()
|
||
if scope in principal.scopes and "research:read" in principal.scopes])
|
||
|
||
async def call_tool(self, ctx, params):
|
||
principal = ctx.request.state.mcp_principal
|
||
return await self.invoke(principal, params.name, params.arguments or {}, str(ctx.request_id or uuid4()))
|
||
|
||
async def invoke(self, principal, name, arguments, request_id=None):
|
||
"""Invoke with a server-authenticated principal; atomic success audit and post-commit wake."""
|
||
started = time.monotonic()
|
||
request_id = request_id or str(uuid4())
|
||
entry = TOOLS.get(name)
|
||
if not entry:
|
||
return tool_result({"error": ResearchError("UNKNOWN_TOOL", "工具不存在").data}, True)
|
||
schema, method, scope, _ = entry
|
||
if "research:read" not in principal.scopes or scope not in principal.scopes:
|
||
raise HTTPException(403, "MCP 令牌缺少所需权限")
|
||
digest = fingerprint(arguments)
|
||
async with self.mutation_lock:
|
||
# Disconnect does not roll back an already accepted operation or lose its wake-up.
|
||
with anyio.CancelScope(shield=True):
|
||
async with self.sessions.begin() as db:
|
||
access = ResearchAccess(db, principal, self.runner.client, self.settings.public_origin)
|
||
code, error = "OK", False
|
||
try:
|
||
async with db.begin_nested():
|
||
args = schema.model_validate(arguments)
|
||
async with asyncio.timeout(30 if method in {"refresh", "metadata", "pyramid_distribution"} else None):
|
||
data = encode_snapshot(await getattr(access, method)(args))
|
||
data.setdefault("_meta", {"schema_version": 1, "observed_at": now().isoformat(),
|
||
"nulls": "null 表示来源未提供,不等于零", "source": "system"})
|
||
except ValidationError as exc:
|
||
code, error = "INVALID_INPUT", True
|
||
data = {"error": ResearchError(code, "; ".join(
|
||
f"{'.'.join(map(str, e['loc']))}: {e['msg']}" for e in exc.errors())).data}
|
||
except TimeoutError:
|
||
code, error = "UPSTREAM_TIMEOUT", True
|
||
data = {"error": ResearchError(code, "元数据读取或刷新超时,未发布新快照", retryable=True).data}
|
||
except ResearchError as exc:
|
||
code, error, data = exc.data["code"], True, {"error": exc.data}
|
||
except HTTPException as exc:
|
||
code = {404: "NOT_FOUND", 409: "CONFLICT", 422: "INVALID_INPUT", 429: "RATE_LIMITED", 502: "UPSTREAM_ERROR"}.get(exc.status_code, "REQUEST_FAILED")
|
||
error = True
|
||
data = {"error": ResearchError(code, str(sanitize(exc.detail)),
|
||
retryable=exc.status_code in {429, 502, 503},
|
||
retry_after=(exc.headers or {}).get("Retry-After")).data}
|
||
except Exception:
|
||
# Never expose SQL parameters, exception reprs or credentials in unexpected errors.
|
||
code, error = "INTERNAL_ERROR", True
|
||
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", data.get("experiment_id", data.get("template_id", data.get("id"))))),
|
||
result_code=code, elapsed_ms=int((time.monotonic()-started)*1000)))
|
||
if not error:
|
||
if access.wake == "backtests":
|
||
self.runner.backtests.wake.set()
|
||
elif access.wake == "jobs":
|
||
self.runner.wake.set()
|
||
return tool_result(data, error)
|