feat: add MCP research access and browser key management
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Authenticated MCP transport; research behavior lives in research_access."""
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Personal access tokens are isolated from browser and upstream credentials."""
|
||||
|
||||
import secrets
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
|
||||
from ..models import Account, Admin, MCPToken, now
|
||||
from ..security import token_hash
|
||||
|
||||
SCOPES = frozenset({"research:read", "research:refresh", "backtests:execute", "backtests:control"})
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Principal:
|
||||
token_id: str
|
||||
admin_id: int
|
||||
account_id: int
|
||||
wq_user_id: str
|
||||
scopes: frozenset[str]
|
||||
|
||||
|
||||
async def create_token(db, name, scopes=None, days=90):
|
||||
"""Issue a token for the bound account; caller commits and reveals it once."""
|
||||
scopes = set(scopes if scopes is not None else ["research:read"])
|
||||
if not name.strip() or len(name) > 100 or not 1 <= days <= 365:
|
||||
raise ValueError("名称须为 1–100 字,有效期须为 1–365 天")
|
||||
if not scopes <= SCOPES or "research:read" not in scopes:
|
||||
raise ValueError("权限无效;所有令牌必须包含 research:read")
|
||||
account, admin = await db.get(Account, 1), await db.get(Admin, 1)
|
||||
if not account or not account.wq_user_id or not admin:
|
||||
raise ValueError("请先初始化系统并确认 WorldQuant 账户身份")
|
||||
secret = "wqmcp_" + secrets.token_urlsafe(32)
|
||||
row = MCPToken(
|
||||
id=str(uuid4()), token_hash=token_hash(secret), name=name.strip(), admin_id=admin.id,
|
||||
account_id=account.id, wq_user_id=account.wq_user_id, scopes=sorted(scopes),
|
||||
expires_at=now() + timedelta(days=days),
|
||||
)
|
||||
db.add(row)
|
||||
await db.flush()
|
||||
return row, secret
|
||||
|
||||
|
||||
async def authenticate(db, secret):
|
||||
"""Validate every request, including current account binding; return no secrets."""
|
||||
row = await db.scalar(select(MCPToken).where(MCPToken.token_hash == token_hash(secret)))
|
||||
if not row or row.revoked_at or row.expires_at.replace(tzinfo=row.expires_at.tzinfo or timezone.utc) <= now():
|
||||
raise HTTPException(401, "MCP 令牌无效或已过期")
|
||||
account, admin = await db.get(Account, row.account_id), await db.get(Admin, row.admin_id)
|
||||
if not account or not admin or account.id != 1 or account.wq_user_id != row.wq_user_id:
|
||||
raise HTTPException(401, "MCP 令牌账户绑定已失效")
|
||||
return Principal(row.id, row.admin_id, row.account_id, row.wq_user_id, frozenset(row.scopes))
|
||||
@@ -0,0 +1,127 @@
|
||||
"""MCP transport over shared research operations, with minimal durable audit evidence."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from uuid import uuid4
|
||||
|
||||
import anyio
|
||||
from fastapi import HTTPException
|
||||
from mcp import types
|
||||
from mcp.server.lowlevel import Server
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from pydantic import ValidationError
|
||||
|
||||
from ..alphas import sanitize
|
||||
from ..backtests.contracts import fingerprint
|
||||
from ..models import MCPAudit, now
|
||||
from ..research.serialization import encode_snapshot
|
||||
from ..research_access import contracts as c
|
||||
from ..research_access.service import ResearchAccess, ResearchError
|
||||
|
||||
# Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
|
||||
TOOLS = {
|
||||
"get_research_capabilities": (c.Empty, "capabilities", "research:read", "读取直接研究能力、完整设置 schema 和调度阻塞,不代表平台剩余额度。"),
|
||||
"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", "查询研究刷新任务的状态和产物引用。"),
|
||||
"search_backtests": (c.History, "history", "research:read", "分页查历史候选与固定设置;candidates 按完整输入精确匹配,不推断数学等价。"),
|
||||
"submit_backtests": (c.Submit, "submit", "backtests:execute", "执行用户已授权的固定批次,自动留痕并立即返回运行 ID。每项必须完整设置;重复默认拒绝,rerun 明确重跑。不需要研究资产。"),
|
||||
"get_backtest": (c.RunReference, "run", "research:read", "读取真实运行进度、提交数量和可选增量事件;受理不等于成功。"),
|
||||
"get_backtest_results": (c.Results, "results", "research:read", "分页读取固定快照指标、全部非通过检查及三层状态;缺失指标不补零。"),
|
||||
"get_backtest_artifact": (c.Artifact, "artifact", "research:read", "分页读取候选脱敏快照的顶层键值或独立采集的 PnL;缺缓存不自动刷新。"),
|
||||
"control_backtest": (c.Control, "control", "backtests:control", "对已授权运行暂停、继续、停止或恢复采集;不远程取消、不重提未知模拟。需要版本和幂等键。"),
|
||||
}
|
||||
|
||||
|
||||
def tool_result(data, error=False):
|
||||
data = encode_snapshot(data)
|
||||
return types.CallToolResult(content=[types.TextContent(type="text", text=json.dumps(data, ensure_ascii=False))],
|
||||
structuredContent=data, isError=error)
|
||||
|
||||
|
||||
class MCPResearchServer:
|
||||
def __init__(self, sessions, runner, settings):
|
||||
self.sessions, self.runner, self.settings = sessions, runner, settings
|
||||
# The existing deployment has one owner; this also gives SQLite test transactions a fair queue.
|
||||
self.mutation_lock = asyncio.Lock()
|
||||
self.server = Server("wq-alpha-research", version="1.0.0", on_list_tools=self.list_tools,
|
||||
on_call_tool=self.call_tool,
|
||||
instructions="自由探索,直接固定候选回测,无需先建研究资产。工具不安排定时研究;结果按运行 ID 查询。")
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
host = urlsplit(settings.public_origin).netloc
|
||||
self.app = self.server.streamable_http_app(
|
||||
streamable_http_path="/", stateless_http=True, json_response=True,
|
||||
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=True,
|
||||
allowed_hosts=[host], allowed_origins=[settings.public_origin.rstrip("/")]),
|
||||
)
|
||||
|
||||
async def list_tools(self, ctx, params):
|
||||
principal = ctx.request.state.mcp_principal
|
||||
return types.ListToolsResult(tools=[types.Tool(name=name, description=description,
|
||||
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"}))
|
||||
for name, (schema, method, scope, description) in TOOLS.items()
|
||||
if scope in principal.scopes and "research:read" in principal.scopes])
|
||||
|
||||
async def call_tool(self, ctx, params):
|
||||
principal = ctx.request.state.mcp_principal
|
||||
return await self.invoke(principal, params.name, params.arguments or {}, str(ctx.request_id or uuid4()))
|
||||
|
||||
async def invoke(self, principal, name, arguments, request_id=None):
|
||||
"""Invoke with a server-authenticated principal; atomic success audit and post-commit wake."""
|
||||
started = time.monotonic()
|
||||
request_id = request_id or str(uuid4())
|
||||
entry = TOOLS.get(name)
|
||||
if not entry:
|
||||
return tool_result({"error": ResearchError("UNKNOWN_TOOL", "工具不存在").data}, True)
|
||||
schema, method, scope, _ = entry
|
||||
if "research:read" not in principal.scopes or scope not in principal.scopes:
|
||||
raise HTTPException(403, "MCP 令牌缺少所需权限")
|
||||
digest = fingerprint(arguments)
|
||||
async with self.mutation_lock:
|
||||
# Disconnect does not roll back an already accepted operation or lose its wake-up.
|
||||
with anyio.CancelScope(shield=True):
|
||||
async with self.sessions.begin() as db:
|
||||
access = ResearchAccess(db, principal, self.runner.client, self.settings.public_origin)
|
||||
code, error = "OK", False
|
||||
try:
|
||||
async with db.begin_nested():
|
||||
args = schema.model_validate(arguments)
|
||||
async with asyncio.timeout(30 if method in {"refresh", "metadata"} else None):
|
||||
data = encode_snapshot(await getattr(access, method)(args))
|
||||
data.setdefault("_meta", {"schema_version": 1, "observed_at": now().isoformat(),
|
||||
"nulls": "null 表示来源未提供,不等于零", "source": "system"})
|
||||
except ValidationError as exc:
|
||||
code, error = "INVALID_INPUT", True
|
||||
data = {"error": ResearchError(code, "; ".join(
|
||||
f"{'.'.join(map(str, e['loc']))}: {e['msg']}" for e in exc.errors())).data}
|
||||
except TimeoutError:
|
||||
code, error = "UPSTREAM_TIMEOUT", True
|
||||
data = {"error": ResearchError(code, "元数据读取或刷新超时,未发布新快照", retryable=True).data}
|
||||
except ResearchError as exc:
|
||||
code, error, data = exc.data["code"], True, {"error": exc.data}
|
||||
except HTTPException as exc:
|
||||
code = {404: "NOT_FOUND", 409: "CONFLICT", 422: "INVALID_INPUT", 429: "RATE_LIMITED", 502: "UPSTREAM_ERROR"}.get(exc.status_code, "REQUEST_FAILED")
|
||||
error = True
|
||||
data = {"error": ResearchError(code, str(sanitize(exc.detail)),
|
||||
retryable=exc.status_code in {429, 502, 503},
|
||||
retry_after=(exc.headers or {}).get("Retry-After")).data}
|
||||
except Exception:
|
||||
# Never expose SQL parameters, exception reprs or credentials in unexpected errors.
|
||||
code, error = "INTERNAL_ERROR", True
|
||||
data = {"error": ResearchError(code, "研究操作失败;可使用原幂等键重试或查询历史", retryable=True).data}
|
||||
db.add(MCPAudit(id=str(uuid4()), token_id=principal.token_id, tool=name,
|
||||
request_id=fingerprint({"request_id": request_id}), input_digest=digest,
|
||||
business_id=data.get("backtest_run_id", data.get("job_id")),
|
||||
result_code=code, elapsed_ms=int((time.monotonic()-started)*1000)))
|
||||
if not error:
|
||||
if access.wake == "backtests":
|
||||
self.runner.backtests.wake.set()
|
||||
elif access.wake == "jobs":
|
||||
self.runner.wake.set()
|
||||
return tool_result(data, error)
|
||||
@@ -0,0 +1,82 @@
|
||||
"""Cookie-authenticated PAT administration; MCP bearer tokens grant no access here."""
|
||||
|
||||
from datetime import timezone
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from ..models import Account, MCPToken, now
|
||||
from ..security import require_auth
|
||||
from .auth import create_token
|
||||
|
||||
router = APIRouter(prefix="/api/v1/mcp-tokens", tags=["mcp-tokens"], dependencies=[Depends(require_auth)])
|
||||
|
||||
|
||||
class TokenInput(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
name: str = Field(min_length=1, max_length=100)
|
||||
days: int = Field(default=90, ge=1, le=365, strict=True)
|
||||
scopes: list[str] = Field(default_factory=lambda: ["research:read"], max_length=4)
|
||||
|
||||
|
||||
def token_output(row, account):
|
||||
"""Return public metadata only, including whether the current binding is usable."""
|
||||
def timestamp(value):
|
||||
return value.replace(tzinfo=value.tzinfo or timezone.utc) if value else None
|
||||
|
||||
expires = timestamp(row.expires_at)
|
||||
status = (
|
||||
"revoked" if row.revoked_at else
|
||||
"expired" if expires <= now() else
|
||||
"invalid_binding" if not account or account.wq_user_id != row.wq_user_id else
|
||||
"active"
|
||||
)
|
||||
return {
|
||||
"id": row.id, "name": row.name, "scopes": row.scopes,
|
||||
"created_at": timestamp(row.created_at), "expires_at": expires,
|
||||
"revoked_at": timestamp(row.revoked_at), "status": status,
|
||||
}
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def list_tokens(request: Request, limit: int = Query(25, ge=1, le=100), offset: int = Query(0, ge=0)):
|
||||
async with request.app.state.sessions() as db:
|
||||
account = await db.get(Account, 1)
|
||||
owned = (MCPToken.admin_id == 1, MCPToken.account_id == 1)
|
||||
total = await db.scalar(select(func.count()).select_from(MCPToken).where(*owned))
|
||||
rows = await db.scalars(select(MCPToken).where(*owned).order_by(
|
||||
MCPToken.created_at.desc(), MCPToken.id.desc()).offset(offset).limit(limit))
|
||||
return {
|
||||
"items": [token_output(row, account) for row in rows], "total": total,
|
||||
"limit": limit, "offset": offset, "has_more": offset + limit < total,
|
||||
"enabled": request.app.state.settings.mcp_enabled,
|
||||
"endpoint": request.app.state.settings.public_origin.rstrip("/") + "/api/v1/mcp/",
|
||||
"can_create": bool(account and account.wq_user_id),
|
||||
}
|
||||
|
||||
|
||||
@router.post("", status_code=201)
|
||||
async def issue_token(body: TokenInput, request: Request):
|
||||
# The existing single-admin browser session is the authority, never request-supplied IDs.
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
try:
|
||||
row, secret = await create_token(db, body.name, body.scopes, body.days)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(422, str(exc)) from exc
|
||||
result = token_output(row, await db.get(Account, 1))
|
||||
# Do not expose the secret until the transaction successfully commits.
|
||||
return {**result, "token": secret}
|
||||
|
||||
|
||||
@router.post("/{token_id}/revoke")
|
||||
async def revoke_token(token_id: str, request: Request):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
row = await db.scalar(select(MCPToken).where(
|
||||
MCPToken.id == token_id, MCPToken.admin_id == 1, MCPToken.account_id == 1,
|
||||
).with_for_update())
|
||||
if not row:
|
||||
raise HTTPException(404, "MCP Key 不存在")
|
||||
row.revoked_at = row.revoked_at or now()
|
||||
result = token_output(row, await db.get(Account, 1))
|
||||
return result
|
||||
Reference in New Issue
Block a user