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
+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 = {
"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"})
+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
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"
+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 ..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(),