feat(mcp): add quarterly pyramid distribution lookup
This commit is contained in:
@@ -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,其余零。此前默认周期样例已被此次明确季度结果替代。
|
||||||
@@ -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"})
|
||||||
|
|||||||
@@ -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()))
|
||||||
@@ -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"
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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(),
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user