feat(mcp): add quarterly pyramid distribution lookup

This commit is contained in:
yuxuanhui
2026-09-13 10:39:31 +08:00
parent db328c62dc
commit 7c8188df9c
8 changed files with 259 additions and 3 deletions
@@ -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,其余零。此前默认周期样例已被此次明确季度结果替代。
+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. # Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
TOOLS = { 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", "分页查询数据准备集合,返回固定范围、字段数及版本。研究可使用多个集合,各集合范围独立。"), "search_data_preparations": (c.PreparationSearch, "preparations", "research:read", "分页查询数据准备集合,返回固定范围、字段数及版本。研究可使用多个集合,各集合范围独立。"),
"get_data_preparation": (c.PreparationRead, "preparation", "research:read", "按集合 ID 与版本分页预览字段、类型、描述和数据集归属。提交回测时携带 preparation_refs,由服务端核对版本并固定独立输入快照;空集合不能用于研究。"), "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、版本和理论组合数;仅核验结构及来源,不验证所有参数组合,不再次调用模型、不执行回测、不覆盖已有模板。"), "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( inputSchema=schema.model_json_schema(), annotations=types.ToolAnnotations(
readOnlyHint=scope == "research:read", destructiveHint=method == "control", readOnlyHint=scope == "research:read", destructiveHint=method == "control",
idempotentHint=method in {"submit", "control", "create_template"} 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"})) openWorldHint=method in {"refresh", "submit", "metadata", "check_self_correlation", "check_submission", "authenticate", "pyramid_distribution"}))
for name, (schema, method, scope, description) in TOOLS.items() for name, (schema, method, scope, description) in TOOLS.items()
if scope in principal.scopes and "research:read" in principal.scopes]) if scope in principal.scopes and "research:read" in principal.scopes])
@@ -101,7 +102,7 @@ class MCPResearchServer:
try: try:
async with db.begin_nested(): async with db.begin_nested():
args = schema.model_validate(arguments) 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 = encode_snapshot(await getattr(access, method)(args))
data.setdefault("_meta", {"schema_version": 1, "observed_at": now().isoformat(), data.setdefault("_meta", {"schema_version": 1, "observed_at": now().isoformat(),
"nulls": "null 表示来源未提供,不等于零", "source": "system"}) "nulls": "null 表示来源未提供,不等于零", "source": "system"})
+80
View File
@@ -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()))
+6
View File
@@ -20,6 +20,12 @@ class Empty(Contract):
pass 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): class Authentication(Contract):
action: Literal["connect", "verify"] = "connect" action: Literal["connect", "verify"] = "connect"
+41
View File
@@ -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
+20
View File
@@ -25,6 +25,7 @@ from ..research.workspace_contracts import FieldAvailabilityInput
from ..schemas import JobInput from ..schemas import JobInput
from ..submission import CheckInput, correlation_allows_check, create_check_job, local_alpha, source from ..submission import CheckInput, correlation_allows_check, create_check_job, local_alpha, source
from ..submission import fingerprint as submission_fingerprint from ..submission import fingerprint as submission_fingerprint
from ..worldquant import WqError
from .contracts import DirectCandidate, History from .contracts import DirectCandidate, History
from .queries import EvidenceQueries, page from .queries import EvidenceQueries, page
@@ -48,6 +49,25 @@ class ResearchAccess:
def run_url(self, run_id): def run_url(self, run_id):
return f"{self.public_origin}/#backtests?run_id={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): async def capabilities(self, args):
return {**await self.backtests.capabilities(), "max_candidates": 100, return {**await self.backtests.capabilities(), "max_candidates": 100,
"settings_schema": DirectCandidate.model_json_schema(), "settings_schema": DirectCandidate.model_json_schema(),
+2 -1
View File
@@ -162,7 +162,8 @@ async def test_official_sdk_client_and_error_contract(mcp_app):
async with ClientSession(streams[0], streams[1]) as client: async with ClientSession(streams[0], streams[1]) as client:
await client.initialize() await client.initialize()
listed = await client.list_tools() 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} assert {"search_data_preparations", "get_data_preparation"} <= {t.name for t in listed.tools}
caps = await client.call_tool("get_research_capabilities", {}) caps = await client.call_tool("get_research_capabilities", {})
assert caps.structured_content["max_candidates"] == 100 assert caps.structured_content["max_candidates"] == 100
+92
View File
@@ -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)