feat: expose local self-correlation tools through MCP
Deploy production / deploy (push) Successful in 57s
Deploy production / deploy (push) Successful in 57s
This commit is contained in:
@@ -25,7 +25,9 @@ TOOLS = {
|
||||
"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", "查询研究刷新任务的状态和产物引用。"),
|
||||
"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", "执行用户已授权的固定批次,自动留痕并立即返回运行 ID。每项必须完整设置;重复默认拒绝,rerun 明确重跑。不需要研究资产。"),
|
||||
"get_backtest": (c.RunReference, "run", "research:read", "读取真实运行进度、提交数量和可选增量事件;受理不等于成功。"),
|
||||
@@ -64,7 +66,7 @@ class MCPResearchServer:
|
||||
inputSchema=schema.model_json_schema(), annotations=types.ToolAnnotations(
|
||||
readOnlyHint=scope == "research:read", destructiveHint=method == "control",
|
||||
idempotentHint=method in {"submit", "control"} or scope == "research:read",
|
||||
openWorldHint=method in {"refresh", "submit", "metadata"}))
|
||||
openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation"}))
|
||||
for name, (schema, method, scope, description) in TOOLS.items()
|
||||
if scope in principal.scopes and "research:read" in principal.scopes])
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from ..schemas import Contract
|
||||
|
||||
Identifier = Annotated[str, Field(min_length=1, max_length=100)]
|
||||
RunId = Annotated[str, Field(min_length=1, max_length=36)]
|
||||
AlphaId = Annotated[str, Field(min_length=1, max_length=100, pattern=r"^[A-Za-z0-9_-]+$")]
|
||||
|
||||
|
||||
class Empty(Contract):
|
||||
@@ -124,6 +125,14 @@ class JobReference(Page):
|
||||
job_id: RunId
|
||||
|
||||
|
||||
class SelfCorrelationCheck(Contract):
|
||||
alpha_ids: list[AlphaId] = Field(min_length=1, max_length=100)
|
||||
|
||||
|
||||
class SelfCorrelationReference(Contract):
|
||||
alpha_id: AlphaId
|
||||
|
||||
|
||||
class History(Page):
|
||||
source: str | None = Field(default=None, max_length=100)
|
||||
reference: str | None = Field(default=None, max_length=200)
|
||||
|
||||
@@ -16,6 +16,7 @@ from ..catalog.contracts import CatalogJobInput
|
||||
from ..catalog.platform import platform_options, validate_platform_scope
|
||||
from ..catalog.research_metadata import ResearchMetadata, availability_key
|
||||
from ..catalog.service import Catalog
|
||||
from ..correlation import MIN_SAMPLES, THRESHOLD, WINDOW_YEARS
|
||||
from ..models import Account, Alpha, BacktestItem, Job, JobItem, ResearchRequest, SimulationAttempt, now
|
||||
from ..research.serialization import encode_snapshot
|
||||
from ..research.workspace_contracts import FieldAvailabilityInput
|
||||
@@ -48,7 +49,15 @@ class ResearchAccess:
|
||||
"settings_schema": DirectCandidate.model_json_schema(),
|
||||
"confirmation": "调用者须已获本批执行授权;直接提交后返回稳定运行 ID",
|
||||
"duplicate_policies": ["reject", "rerun"], "permissions": sorted(self.principal.scopes),
|
||||
"metadata_only": True, "actual_platform_allowance": None}
|
||||
"metadata_only": True, "actual_platform_allowance": None,
|
||||
"self_correlation": {
|
||||
"check_with": "check_self_correlation", "read_with": "get_self_correlation",
|
||||
"job_with": "get_refresh_job", "max_targets": 100, "source": "local",
|
||||
"reference_scope": "本地已同步的同地区已提交 Alpha,排除自身",
|
||||
"method": "累计 PnL 日变化的 Pearson 相关系数,取带符号最大值",
|
||||
"threshold": THRESHOLD, "min_samples": MIN_SAMPLES, "window_years": WINDOW_YEARS,
|
||||
"platform_check": False,
|
||||
}}
|
||||
|
||||
async def catalog(self, args):
|
||||
data = await Catalog(self.db).search(args.filters, args.dataset_id)
|
||||
@@ -101,7 +110,7 @@ class ResearchAccess:
|
||||
|
||||
async def refresh_job(self, args):
|
||||
job = await self.db.get(Job, args.job_id)
|
||||
if not job or job.kind not in {"catalog_sync", "field_sync", "pnl_refresh"}:
|
||||
if not job or job.kind not in {"catalog_sync", "field_sync", "pnl_refresh", "self_correlation"}:
|
||||
raise ResearchError("NOT_FOUND", "研究刷新任务不存在")
|
||||
result = await self.business.get_job_status(args.job_id)
|
||||
query = select(JobItem).where(JobItem.job_id == job.id, JobItem.error.is_not(None))
|
||||
@@ -111,6 +120,27 @@ class ResearchAccess:
|
||||
return {**result, "job_id": job.id, "artifact_reference": job.payload,
|
||||
"errors": page([{"alpha_id": e.alpha_id, "error": e.error} for e in errors], total, args.limit, args.offset)}
|
||||
|
||||
async def check_self_correlation(self, args):
|
||||
"""Queue local checks for synced IDs; caller commits before waking the runner.
|
||||
|
||||
The shared job service deduplicates active batches. Missing PnL is fetched
|
||||
by the durable runner, so slow upstream reads do not hold the MCP call.
|
||||
"""
|
||||
ids = sorted(set(args.alpha_ids))
|
||||
existing = set(await self.db.scalars(select(Alpha.id).where(Alpha.id.in_(ids))))
|
||||
if existing != set(ids):
|
||||
raise ResearchError("NOT_FOUND", "部分 Alpha 尚未同步,请先导入", affected_items=sorted(set(ids)-existing))
|
||||
job = await self.business.create_sync_job(JobInput(kind="self_correlation", alpha_ids=ids))
|
||||
self.wake = "jobs"
|
||||
return {"job_id": job["id"], "status": job["status"], "alpha_ids": ids,
|
||||
"source": "local", "job_with": "get_refresh_job", "read_with": "get_self_correlation"}
|
||||
|
||||
async def self_correlation(self, args):
|
||||
"""Read the latest local result without fetching PnL or starting a check."""
|
||||
data = await self.business.get_self_correlation(args.alpha_id)
|
||||
status = "not_cached" if not data["cached"] else "stale" if data["result"]["stale"] else "available"
|
||||
return {"alpha_id": args.alpha_id, "source": "local", "status": status, **data}
|
||||
|
||||
async def history(self, args):
|
||||
return await self.evidence.history(args)
|
||||
|
||||
|
||||
+106
-1
@@ -162,7 +162,7 @@ async def test_official_sdk_client_and_error_contract(mcp_app):
|
||||
async with ClientSession(streams[0], streams[1]) as client:
|
||||
await client.initialize()
|
||||
listed = await client.list_tools()
|
||||
assert len(listed.tools) == 11
|
||||
assert len(listed.tools) == 13
|
||||
caps = await client.call_tool("get_research_capabilities", {})
|
||||
assert caps.structured_content["max_candidates"] == 100
|
||||
result = await client.call_tool("submit_backtests", submission())
|
||||
@@ -278,3 +278,108 @@ async def test_refresh_job_error_pages_and_kind_isolation(mcp_app):
|
||||
assert result["errors"]["items"][0]["alpha_id"] == "1"
|
||||
error = await mcp_app.state.mcp.invoke(principal, "get_refresh_job", {"job_id": "auth"})
|
||||
assert error.structured_content["error"]["code"] == "NOT_FOUND"
|
||||
|
||||
|
||||
async def test_self_correlation_sdk_workflow_cache_and_staleness(mcp_app, monkeypatch):
|
||||
import httpx2
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
|
||||
from app.alphas import upsert_alpha
|
||||
from app.models import Job, SelfCorrelation
|
||||
from tests.conftest import alpha
|
||||
from tests.test_alpha_management import points
|
||||
|
||||
data = points([1, -2, 4, 3] * 20)
|
||||
async with mcp_app.state.sessions.begin() as db:
|
||||
for raw in (alpha("target"), alpha("peer", status="ACTIVE"), alpha("pending")):
|
||||
await upsert_alpha(db, raw)
|
||||
db.add(Pnl(alpha_id="target", raw={}, points=data))
|
||||
calls = []
|
||||
|
||||
async def pnl(alpha_id):
|
||||
calls.append(alpha_id)
|
||||
return {"records": [{"date": p["date"], "pnl": p["value"]} for p in data]}
|
||||
|
||||
monkeypatch.setattr(mcp_app.state.runner.client, "pnl", pnl)
|
||||
principal, secret = await credentials(mcp_app, {"research:read", "research:refresh"})
|
||||
async with httpx2.AsyncClient(transport=httpx2.ASGITransport(app=mcp_app),
|
||||
headers={"Authorization": f"Bearer {secret}"}) as http:
|
||||
async with streamable_http_client("http://testserver/api/v1/mcp/", http_client=http) as streams:
|
||||
async with ClientSession(streams[0], streams[1]) as client:
|
||||
await client.initialize()
|
||||
listed = {t.name: t for t in (await client.list_tools()).tools}
|
||||
assert listed["get_self_correlation"].annotations.read_only_hint
|
||||
assert not listed["check_self_correlation"].annotations.read_only_hint
|
||||
assert listed["check_self_correlation"].annotations.open_world_hint
|
||||
before = await client.call_tool("get_self_correlation", {"alpha_id": "target"})
|
||||
assert before.structured_content["status"] == "not_cached"
|
||||
assert before.structured_content["result"] is None and calls == []
|
||||
mcp_app.state.runner.wake.clear()
|
||||
started = await client.call_tool("check_self_correlation", {"alpha_ids": ["target", "target"]})
|
||||
assert not started.is_error, started
|
||||
job_id = started.structured_content["job_id"]
|
||||
assert mcp_app.state.runner.wake.is_set() and calls == []
|
||||
assert started.structured_content["alpha_ids"] == ["target"]
|
||||
again = await client.call_tool("check_self_correlation", {"alpha_ids": ["target"]})
|
||||
assert again.structured_content["job_id"] == job_id
|
||||
queued = await client.call_tool("get_refresh_job", {"job_id": job_id})
|
||||
assert queued.structured_content["status"] == "queued"
|
||||
await mcp_app.state.runner.execute(job_id)
|
||||
done = await client.call_tool("get_refresh_job", {"job_id": job_id})
|
||||
assert done.structured_content["status"] == "completed"
|
||||
assert done.structured_content["processed"] == 1
|
||||
result = await client.call_tool("get_self_correlation", {"alpha_id": "target"})
|
||||
body = result.structured_content
|
||||
assert body["source"] == "local" and body["status"] == "available"
|
||||
assert body["result"]["max_correlation"] == pytest.approx(1)
|
||||
assert body["result"]["compared_count"] == 1
|
||||
assert body["result"]["matches"][0]["alpha_id"] == "peer"
|
||||
assert calls == ["peer"]
|
||||
async with mcp_app.state.sessions.begin() as db:
|
||||
row = await db.get(SelfCorrelation, "target")
|
||||
row.stale = True
|
||||
audit = await db.scalar(select(MCPAudit).where(MCPAudit.tool == "check_self_correlation"))
|
||||
assert audit.business_id == job_id and audit.result_code == "OK"
|
||||
assert await db.scalar(select(func.count()).select_from(Job).where(Job.kind == "self_correlation")) == 1
|
||||
assert await db.get(Pnl, "peer") is not None
|
||||
stale = await invoke(mcp_app, principal, "get_self_correlation", {"alpha_id": "target"})
|
||||
assert stale["status"] == "stale" and stale["result"]["stale"] and calls == ["peer"]
|
||||
|
||||
|
||||
async def test_self_correlation_read_only_scope_and_invalid_inputs(mcp_app):
|
||||
from app.alphas import upsert_alpha
|
||||
from app.models import Job
|
||||
from tests.conftest import alpha
|
||||
|
||||
async with mcp_app.state.sessions.begin() as db:
|
||||
await upsert_alpha(db, alpha("known"))
|
||||
_, secret = await credentials(mcp_app, {"research:read"})
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=mcp_app), base_url="http://testserver",
|
||||
headers={"Authorization": f"Bearer {secret}", "Accept": "application/json, text/event-stream"}) as http:
|
||||
listed = await http.post(ENDPOINT, json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"})
|
||||
names = {t["name"] for t in listed.json()["result"]["tools"]}
|
||||
assert "get_self_correlation" in names and "check_self_correlation" not in names
|
||||
for name, args, status in (
|
||||
("check_self_correlation", {"alpha_ids": ["known"]}, 403),
|
||||
("get_self_correlation", {"alpha_id": "known"}, 200),
|
||||
):
|
||||
response = await http.post(ENDPOINT, json={"jsonrpc": "2.0", "id": 2, "method": "tools/call",
|
||||
"params": {"name": name, "arguments": args}})
|
||||
assert response.status_code == status
|
||||
principal, _ = await credentials(mcp_app)
|
||||
for args, code in (
|
||||
({"alpha_ids": []}, "INVALID_INPUT"),
|
||||
({"alpha_ids": ["known"] * 101}, "INVALID_INPUT"),
|
||||
({"alpha_ids": ["../secret"]}, "INVALID_INPUT"),
|
||||
({"alpha_ids": ["known"], "force": True}, "INVALID_INPUT"),
|
||||
({"alpha_ids": ["known", "missing"]}, "NOT_FOUND"),
|
||||
):
|
||||
failure = await mcp_app.state.mcp.invoke(principal, "check_self_correlation", args)
|
||||
assert failure.is_error and failure.structured_content["error"]["code"] == code
|
||||
if code == "NOT_FOUND":
|
||||
assert failure.structured_content["error"]["affected_items"] == ["missing"]
|
||||
missing = await mcp_app.state.mcp.invoke(principal, "get_self_correlation", {"alpha_id": "missing"})
|
||||
assert missing.is_error and missing.structured_content["error"]["code"] == "NOT_FOUND"
|
||||
async with mcp_app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(Job)) == 0
|
||||
|
||||
Reference in New Issue
Block a user