From 7c8188df9c95b08aa6c63d29d52828816f78803b Mon Sep 17 00:00:00 2001 From: yuxuanhui Date: Sun, 13 Sep 2026 10:39:31 +0800 Subject: [PATCH] feat(mcp): add quarterly pyramid distribution lookup --- .../pyramid-mcp/issues/01-distribution.md | 15 +++ backend/app/mcp_api/server.py | 5 +- backend/app/probe_pyramids.py | 80 ++++++++++++++++ backend/app/research_access/contracts.py | 6 ++ backend/app/research_access/pyramids.py | 41 +++++++++ backend/app/research_access/service.py | 20 ++++ backend/tests/test_mcp.py | 3 +- backend/tests/test_mcp_pyramids.py | 92 +++++++++++++++++++ 8 files changed, 259 insertions(+), 3 deletions(-) create mode 100644 .scratch/pyramid-mcp/issues/01-distribution.md create mode 100644 backend/app/probe_pyramids.py create mode 100644 backend/app/research_access/pyramids.py create mode 100644 backend/tests/test_mcp_pyramids.py diff --git a/.scratch/pyramid-mcp/issues/01-distribution.md b/.scratch/pyramid-mcp/issues/01-distribution.md new file mode 100644 index 0000000..5f1d3ab --- /dev/null +++ b/.scratch/pyramid-mcp/issues/01-distribution.md @@ -0,0 +1,15 @@ +# Pyramid distribution MCP +Status: ready-for-agent + +按 region/delay 实时读取本季度个人 Pyramid 分布,按 >=3、1–2、0 分组。 +复用平台认证和 MCP 只读权限;缺失或非法计数不能当作零。 +验证 MCP 调用、边界值、输入/上游错误,并展示实时样例。不发布、不提交。 + +## Comments + +- 已实现 get_pyramid_distribution(region, delay),复用现有 WqClient 和 research:read。 +- 实测传日期后返回全零,日期参数语义尚未确认;最终使用已验证的无日期请求,period=platform_default,不宣称独立验证季度边界。 +- MCP 及新功能 20 项测试通过,Ruff / diff check 通过。使用内存数据库隔离认证令牌与审计,通过 MCP invoke 调用真实平台:USA/D1 已点亮0类、进行中5类、零计数11类。 + +- 按用户新要求改为必传 current_date,自动计算完整自然季度,发送 startDate/endDate;不再使用默认周期。2026-09-13 对应 2026-07-01 至 2026-09-30。 +- 33 项测试通过(含四季度边界、跨年、闰日、日期校验),Ruff 通过。真实 MCP invoke:USA/D1 全零;GLB/D1 fundamental/risk 各3;AMR/D1 risk 为1,其余零。此前默认周期样例已被此次明确季度结果替代。 diff --git a/backend/app/mcp_api/server.py b/backend/app/mcp_api/server.py index a72c4be..b332146 100644 --- a/backend/app/mcp_api/server.py +++ b/backend/app/mcp_api/server.py @@ -21,6 +21,7 @@ from ..research_access.service import ResearchAccess, ResearchError # Name, schema, business method, required scope, description. No generic arbitrary HTTP tool. TOOLS = { + "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,由服务端核对版本并固定独立输入快照;空集合不能用于研究。"), "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、版本和理论组合数;仅核验结构及来源,不验证所有参数组合,不再次调用模型、不执行回测、不覆盖已有模板。"), @@ -73,7 +74,7 @@ class MCPResearchServer: inputSchema=schema.model_json_schema(), annotations=types.ToolAnnotations( readOnlyHint=scope == "research:read", destructiveHint=method == "control", idempotentHint=method in {"submit", "control", "create_template"} or scope == "research:read", - openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation", "check_submission", "authenticate"})) + openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation", "check_submission", "authenticate", "pyramid_distribution"})) for name, (schema, method, scope, description) in TOOLS.items() if scope in principal.scopes and "research:read" in principal.scopes]) @@ -101,7 +102,7 @@ class MCPResearchServer: try: async with db.begin_nested(): args = schema.model_validate(arguments) - async with asyncio.timeout(30 if method in {"refresh", "metadata"} else None): + 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"}) diff --git a/backend/app/probe_pyramids.py b/backend/app/probe_pyramids.py new file mode 100644 index 0000000..79109a1 --- /dev/null +++ b/backend/app/probe_pyramids.py @@ -0,0 +1,80 @@ +"""Read-only Pyramid endpoint probe: python -m app.probe_pyramids. + +Uses the project's configured account and WqClient without changing stored data. +A standalone process authenticates separately; it cannot inherit a running +server's in-memory cookies. Authentication responses and secrets are never printed. +""" + +import asyncio +import json + +from .config import Settings +from .db import create_database +from .models import Account +from .security import cipher +from .worldquant import WqClient, WqError + +PATHS = ( + "/users/self/activities/pyramid-alphas", + "/users/self/pyramid/alphas", + "/activities/pyramid-alphas", + "/pyramid/alphas", +) + + +async def probe(client): + """Probe fixed same-origin GET paths using an authenticated WqClient. + + Returns status and successful JSON for each path; stops on session expiry. + Transport failures propagate to the caller without exposing request details. + """ + if not client.authenticated: + raise WqError("请先连接 WorldQuant", "disconnected") + results = [] + for path in PATHS: + response = await client.client.get(path) + result = {"path": path, "status": response.status_code} + if response.status_code == 200: + try: + result["data"] = response.json() + except ValueError: + result["error"] = "invalid_json" + results.append(result) + if response.status_code in (401, 429): + break + return results + + +async def main(): + settings = Settings() + client = WqClient(settings) + engine = None + try: + if settings.wq_email: + email, password = settings.wq_email, settings.wq_password.get_secret_value() + else: + engine, sessions = create_database(settings.database_url) + async with sessions() as db: + account = await db.get(Account, 1) + if not account or not account.email or not account.password_encrypted: + raise WqError("未配置平台凭据", "disconnected") + email = account.email + password = cipher(settings).decrypt(account.password_encrypted.encode()).decode() + await client.authenticate(email, password) + print(json.dumps(await probe(client), ensure_ascii=False, indent=2)) + except WqError as exc: + print(json.dumps({"error": exc.code}, ensure_ascii=False)) + return 1 + except Exception as exc: + # Connection/config errors can contain credentials; print only the type. + print(json.dumps({"error": type(exc).__name__})) + return 1 + finally: + await client.close() + if engine is not None: + await engine.dispose() + return 0 + + +if __name__ == "__main__": + raise SystemExit(asyncio.run(main())) diff --git a/backend/app/research_access/contracts.py b/backend/app/research_access/contracts.py index 3a7d34a..a8443aa 100644 --- a/backend/app/research_access/contracts.py +++ b/backend/app/research_access/contracts.py @@ -20,6 +20,12 @@ class Empty(Contract): pass +class PyramidQuery(Contract): + current_date: date = Field(description="用于确定季度的日期,格式 YYYY-MM-DD;自动查询该季度完整起止范围") + region: str = Field(min_length=3, max_length=10, pattern=r"^[A-Z]+$") + delay: int = Field(ge=0, le=1, strict=True) + + class Authentication(Contract): action: Literal["connect", "verify"] = "connect" diff --git a/backend/app/research_access/pyramids.py b/backend/app/research_access/pyramids.py new file mode 100644 index 0000000..f402d0f --- /dev/null +++ b/backend/app/research_access/pyramids.py @@ -0,0 +1,41 @@ +"""Partition platform category counts without treating missing evidence as zero.""" + +from calendar import monthrange + + +def quarter_period(current_date): + """Return the full calendar quarter containing the supplied date, inclusive.""" + quarter = (current_date.month - 1) // 3 + 1 + end_month = quarter * 3 + start = current_date.replace(month=end_month - 2, day=1) + end = current_date.replace(month=end_month, day=monthrange(current_date.year, end_month)[1]) + return {"quarter": f"{current_date.year}-Q{quarter}", + "start_date": start.isoformat(), "end_date": end.isoformat()} + + +def distribution(raw, region, delay): + """Return three category lists; raise ValueError on incomplete or duplicate data.""" + if not isinstance(raw, dict) or not isinstance(raw.get("pyramids"), list): + raise ValueError("平台未提供 Pyramid 分布") + groups = {"lit": [], "in_progress": [], "unlit": []} + seen = set() + for row in raw["pyramids"]: + if not isinstance(row, dict): + raise ValueError("平台 Pyramid 数据格式异常") + if row.get("region") != region or row.get("delay") != delay: + continue + category, count = row.get("category"), row.get("alphaCount") + if (not isinstance(category, dict) + or not isinstance(category.get("id"), str) or not category["id"] + or not isinstance(category.get("name"), str) or not category["name"] + or type(count) is not int or count < 0 or category["id"] in seen): + raise ValueError("平台分类或计数缺失、非法或重复,不能判定点塔状态") + seen.add(category["id"]) + key = "lit" if count >= 3 else "in_progress" if count > 0 else "unlit" + groups[key].append({"category": {"id": category["id"], "name": category["name"]}, + "alpha_count": count, "remaining": max(0, 3 - count)}) + if not seen: + raise ValueError("平台未返回此 region/delay 的分类,不能认定全部未点亮") + for items in groups.values(): + items.sort(key=lambda item: item["category"]["id"]) + return groups diff --git a/backend/app/research_access/service.py b/backend/app/research_access/service.py index 0ef14b8..5f75ec7 100644 --- a/backend/app/research_access/service.py +++ b/backend/app/research_access/service.py @@ -25,6 +25,7 @@ from ..research.workspace_contracts import FieldAvailabilityInput from ..schemas import JobInput from ..submission import CheckInput, correlation_allows_check, create_check_job, local_alpha, source from ..submission import fingerprint as submission_fingerprint +from ..worldquant import WqError from .contracts import DirectCandidate, History from .queries import EvidenceQueries, page @@ -48,6 +49,25 @@ class ResearchAccess: def run_url(self, run_id): return f"{self.public_origin}/#backtests?run_id={run_id}" + async def pyramid_distribution(self, args): + """Read the supplied date's full quarter using the three-Alpha completion rule.""" + from .pyramids import distribution, quarter_period + + account = await self.db.get(Account, self.principal.account_id) + if not account or account.wq_user_id != self.principal.wq_user_id: + raise ResearchError("ACCOUNT_MISMATCH", "平台账户绑定已变化") + period = quarter_period(args.current_date) + try: + raw = await self.client.get_pyramid_alphas(period["start_date"], period["end_date"]) + groups = distribution(raw, args.region, args.delay) + except WqError as exc: + raise ResearchError(exc.code.upper(), str(exc), retryable=True) from None + except ValueError as exc: + raise ResearchError("INVALID_PLATFORM_DATA", str(exc)) from None + return {"region": args.region, "delay": args.delay, "threshold": 3, + "current_date": args.current_date.isoformat(), "period": period, + "source": "worldquant_platform", **groups} + async def capabilities(self, args): return {**await self.backtests.capabilities(), "max_candidates": 100, "settings_schema": DirectCandidate.model_json_schema(), diff --git a/backend/tests/test_mcp.py b/backend/tests/test_mcp.py index d0da7e0..b00c7ee 100644 --- a/backend/tests/test_mcp.py +++ b/backend/tests/test_mcp.py @@ -162,7 +162,8 @@ 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) == 20 + assert len(listed.tools) == 21 + assert any(tool.name == "get_pyramid_distribution" for tool in listed.tools) assert {"search_data_preparations", "get_data_preparation"} <= {t.name for t in listed.tools} caps = await client.call_tool("get_research_capabilities", {}) assert caps.structured_content["max_candidates"] == 100 diff --git a/backend/tests/test_mcp_pyramids.py b/backend/tests/test_mcp_pyramids.py new file mode 100644 index 0000000..fa4355c --- /dev/null +++ b/backend/tests/test_mcp_pyramids.py @@ -0,0 +1,92 @@ +"""Verify distribution through MCP authorization, validation and audit boundaries.""" + +import httpx +import pytest + +from app.worldquant import WqClient +from tests.test_mcp import credentials, invoke +from tests.test_mcp import mcp_app as mcp_app + + +def row(name, count, region="USA", delay=1): + return {"category": {"id": name, "name": name}, "alphaCount": count, + "region": region, "delay": delay} + + +async def test_distribution(mcp_app): + calls = [] + + def platform(request): + calls.append(request) + assert request.method == "GET" + assert request.url.path == "/users/self/activities/pyramid-alphas" + return httpx.Response(200, json={"pyramids": [row("zero", 0), row("one", 1), + row("two", 2), row("three", 3), row("four", 4), row("other", 5, "GLB"), + row("delay_zero", 5, delay=0)]}) + + client = WqClient(mcp_app.state.settings, transport=httpx.MockTransport(platform)) + client.credentials, client.authenticated = ("test", "test"), True + original = mcp_app.state.runner.client + mcp_app.state.runner.client = client + try: + reader, _ = await credentials(mcp_app, {"research:read"}) + result = await invoke(mcp_app, reader, "get_pyramid_distribution", {"region": "USA", "delay": 1, "current_date": "2026-09-13"}) + assert [r["alpha_count"] for r in result["lit"]] == [4, 3] + assert [r["remaining"] for r in result["in_progress"]] == [2, 1] + assert result["unlit"][0]["category"]["id"] == "zero" + assert dict(calls[0].url.params) == {"startDate": "2026-07-01", "endDate": "2026-09-30"} + assert result["period"] == {"quarter": "2026-Q3", "start_date": "2026-07-01", "end_date": "2026-09-30"} + for args in ({"region": "USA", "delay": 2}, {"region": "USA", "delay": True}, + {"region": "../", "delay": 1}): + response = await mcp_app.state.mcp.invoke(reader, "get_pyramid_distribution", args | {"current_date": "2026-09-13"}) + assert response.structured_content["error"]["code"] == "INVALID_INPUT" + assert len(calls) == 1 + client.credentials, client.authenticated = None, False + response = await mcp_app.state.mcp.invoke(reader, "get_pyramid_distribution", {"region": "USA", "delay": 1, "current_date": "2026-09-13"}) + assert response.structured_content["error"]["code"] == "DISCONNECTED" + finally: + mcp_app.state.runner.client = original + await client.close() + + +@pytest.mark.parametrize("rows", [[], [row("x", None)], [row("x", -1)], + [row("x", True)], [row("x", 1), row("x", 2)]]) +def test_missing_evidence_is_not_zero(rows): + from app.research_access.pyramids import distribution + + with pytest.raises(ValueError): + distribution({"pyramids": rows}, "USA", 1) + + +@pytest.mark.parametrize(('value', 'quarter', 'start', 'end'), [ + ('2026-01-01', '2026-Q1', '2026-01-01', '2026-03-31'), + ('2026-03-31', '2026-Q1', '2026-01-01', '2026-03-31'), + ('2026-04-01', '2026-Q2', '2026-04-01', '2026-06-30'), + ('2026-06-30', '2026-Q2', '2026-04-01', '2026-06-30'), + ('2026-07-01', '2026-Q3', '2026-07-01', '2026-09-30'), + ('2026-09-30', '2026-Q3', '2026-07-01', '2026-09-30'), + ('2026-10-01', '2026-Q4', '2026-10-01', '2026-12-31'), + ('2026-12-31', '2026-Q4', '2026-10-01', '2026-12-31'), + ('2027-01-01', '2027-Q1', '2027-01-01', '2027-03-31'), + ('2024-02-29', '2024-Q1', '2024-01-01', '2024-03-31'), +]) +def test_quarter_boundaries(value, quarter, start, end): + from datetime import date + + from app.research_access.pyramids import quarter_period + + assert quarter_period(date.fromisoformat(value)) == { + 'quarter': quarter, 'start_date': start, 'end_date': end} + + +@pytest.mark.parametrize('value', [None, '2026-02-30', 'not-a-date']) +def test_date_required_and_valid(value): + from pydantic import ValidationError + + from app.research_access.contracts import PyramidQuery + + args = {'region': 'USA', 'delay': 1} + if value is not None: + args['current_date'] = value + with pytest.raises(ValidationError): + PyramidQuery.model_validate(args)