feat: add MCP research access and browser key management

This commit is contained in:
yuxuanhui
2026-09-09 16:18:27 +08:00
parent 4debca7dbd
commit 45238280e3
47 changed files with 2642 additions and 44 deletions
+1
View File
@@ -0,0 +1 @@
"""Direct research interface shared by trusted application adapters."""
+170
View File
@@ -0,0 +1,170 @@
"""Bounded direct research inputs; unknown properties are rejected at the interface."""
from datetime import date, datetime
from typing import Annotated, Literal
from pydantic import Field, model_validator
from ..backtests.contracts import Candidate, SimulationSettings
from ..catalog.contracts import CatalogFilters, Scope
from ..schemas import Contract
Identifier = Annotated[str, Field(min_length=1, max_length=100)]
RunId = Annotated[str, Field(min_length=1, max_length=36)]
class Empty(Contract):
pass
class Page(Contract):
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class CompleteSettings(SimulationSettings):
@model_validator(mode="before")
@classmethod
def complete(cls, value):
if isinstance(value, dict) and set(cls.model_fields) - value.keys():
raise ValueError("必须提供每项完整设置;先读取 get_research_capabilities")
return value
model_config = {"json_schema_extra": {"required": list(SimulationSettings.model_fields)}}
class DirectCandidate(Candidate):
settings: CompleteSettings
class Provenance(Contract):
reference: str | None = Field(default=None, max_length=200)
batch_id: str | None = Field(default=None, max_length=200)
hypothesis: str | None = Field(default=None, max_length=2000)
parent_run_id: RunId | None = None
class Submit(Contract):
name: str = Field(min_length=1, max_length=200)
candidates: list[DirectCandidate] = Field(min_length=1, max_length=100)
idempotency_key: Identifier
duplicate_policy: Literal["reject", "rerun"] = "reject"
source: Provenance = Field(default_factory=Provenance)
@model_validator(mode="after")
def unique_ids(self):
if len({c.client_item_id for c in self.candidates}) != len(self.candidates):
raise ValueError("client_item_id 必须唯一")
return self
class Control(Contract):
run_id: RunId
action: Literal["pause", "resume", "stop", "recover"]
expected_version: int = Field(ge=1)
idempotency_key: Identifier
class CatalogSearch(Contract):
filters: CatalogFilters
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
class Scopes(Contract):
kind: Literal["scopes"]
class SettingOptions(Page):
kind: Literal["settings"]
class Operators(Page):
kind: Literal["operators"]
q: str = Field(default="", max_length=300)
category: str | None = None
class Availability(Contract):
kind: Literal["field_availability"]
field_id: Identifier
scope: Scope
class Metadata(Contract):
query: Annotated[Scopes | SettingOptions | Operators | Availability, Field(discriminator="kind")]
class CatalogRefresh(Contract):
kind: Literal["catalog"]
scope: Scope
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
class OperatorsRefresh(Contract):
kind: Literal["operators"]
class SettingsRefresh(Contract):
kind: Literal["settings"]
class PnlRefresh(Contract):
kind: Literal["pnl"]
alpha_ids: list[Identifier] = Field(min_length=1, max_length=100)
class Refresh(Contract):
query: Annotated[
CatalogRefresh | OperatorsRefresh | SettingsRefresh | Availability | PnlRefresh,
Field(discriminator="kind"),
]
class JobReference(Page):
job_id: RunId
class History(Page):
source: str | None = Field(default=None, max_length=100)
reference: str | None = Field(default=None, max_length=200)
status: str | None = Field(default=None, max_length=30)
created_from: datetime | None = None
created_to: datetime | None = None
scope: Scope | None = None
q: str = Field(default="", max_length=300)
candidates: list[DirectCandidate] | None = Field(default=None, min_length=1, max_length=100)
@model_validator(mode="after")
def dates(self):
for value in (self.created_from, self.created_to):
if value and not value.tzinfo:
raise ValueError("时间须包含时区")
if self.created_from and self.created_to and self.created_from > self.created_to:
raise ValueError("起始时间不能晚于结束时间")
return self
class RunReference(Contract):
run_id: RunId
after: int | None = Field(default=None, ge=0)
event_limit: int = Field(default=25, ge=1, le=100)
class Results(Page):
run_id: RunId
item_ids: list[RunId] | None = Field(default=None, min_length=1, max_length=100)
class Artifact(Page):
item_id: RunId
kind: Literal["snapshot", "pnl"]
date_from: date | None = None
date_to: date | None = None
@model_validator(mode="after")
def dates(self):
if self.kind == "snapshot" and (self.date_from or self.date_to):
raise ValueError("日期筛选仅用于 PnL")
if self.date_from and self.date_to and self.date_from > self.date_to:
raise ValueError("起始日期不能晚于结束日期")
return self
+134
View File
@@ -0,0 +1,134 @@
"""Historical evidence reads, independent of transport and current Alpha refreshes."""
from collections import Counter
from sqlalchemy import func, select
from ..alphas import number, sanitize
from ..backtests.contracts import fingerprint
from ..models import BacktestItem, BacktestResult, BacktestRun, Pnl
from ..research.serialization import encode_snapshot
def page(items, total, limit, offset):
return {"items": items, "total": total, "limit": limit, "offset": offset,
"has_more": offset + len(items) < total}
def checks_summary(snapshot):
"""Preserve unknown check values; missing checks can never mean passed."""
checks = []
for section in ("is", "os"):
metrics = snapshot.get(section)
if isinstance(metrics, dict) and "checks" in metrics:
raw = metrics["checks"]
checks.extend({"section": section, "raw": c} for c in (raw if isinstance(raw, list) else [raw]))
if "checks" in snapshot:
raw = snapshot["checks"]
checks.extend({"section": "root", "raw": c} for c in (raw if isinstance(raw, list) else [raw]))
counts = Counter({key: 0 for key in ("PASS", "FAIL", "PENDING", "WARNING", "UNKNOWN")})
non_pass = []
for check in checks:
raw = check["raw"]
value = raw.get("result", raw.get("status")) if isinstance(raw, dict) else None
state = value if isinstance(value, str) and value in counts else "UNKNOWN"
counts[state] += 1
if state != "PASS":
non_pass.append({**check, "status": state})
return {"status": "unknown" if not checks else "reported", "counts": dict(counts),
"total": len(checks), "non_pass": non_pass}
def item_summary(item, result):
snapshot = sanitize(result.snapshot) if result else {}
metrics = {}
for section in ("is", "os"):
raw = snapshot.get(section)
raw = raw if isinstance(raw, dict) else {}
metrics[section] = {key: number(raw.get(key)) for key in
("sharpe", "fitness", "returns", "turnover", "margin", "drawdown")}
return encode_snapshot({
**{k: getattr(item, k) for k in (
"id", "run_id", "client_item_id", "expression", "settings", "attempt_id",
"platform_status", "collection_status", "persistence_status", "simulation_id", "alpha_id",
)},
"error": sanitize(item.error), "metrics": metrics,
"missing_metrics_reason": "来源未提供或非有限数字;null 不等于零",
"checks": checks_summary(snapshot),
"result": {"observed_at": result.observed_at, "complete": result.complete} if result else None,
"artifact_reference": {"item_id": item.id},
})
class EvidenceQueries:
def __init__(self, db):
self.db = db
async def history(self, args):
query = select(BacktestItem, BacktestResult, BacktestRun).join(
BacktestRun, BacktestRun.id == BacktestItem.run_id
).outerjoin(BacktestResult, BacktestResult.item_id == BacktestItem.id)
for key in ("source", "reference"):
value = getattr(args, key)
if value is not None:
query = query.where(BacktestRun.source["kind" if key == "source" else key].as_string() == value)
if args.status:
query = query.where(BacktestRun.status == args.status)
if args.created_from:
query = query.where(BacktestRun.created_at >= args.created_from)
if args.created_to:
query = query.where(BacktestRun.created_at <= args.created_to)
if args.scope:
for source, target in (("instrument_type", "instrumentType"), ("region", "region"), ("universe", "universe")):
query = query.where(BacktestItem.settings[target].as_string() == getattr(args.scope, source))
query = query.where(BacktestItem.settings["delay"].as_integer() == args.scope.delay)
if args.q:
escaped = args.q.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
query = query.where(BacktestItem.expression.ilike(f"%{escaped}%", escape="\\"))
matches = {}
if args.candidates:
for c in args.candidates:
matches.setdefault(fingerprint(c.platform_input()), []).append(c.client_item_id)
query = query.where(BacktestItem.fingerprint.in_(matches))
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = (await self.db.execute(query.order_by(BacktestRun.created_at.desc(), BacktestRun.id, BacktestItem.ordinal)
.limit(args.limit).offset(args.offset))).all()
items = [{**item_summary(i, r), "source": run.source, "run_status": run.status,
"created_at": run.created_at, "matched_candidates": matches.get(i.fingerprint, []),
"match_type": "exact_input" if args.candidates else "filter"} for i, r, run in rows]
return encode_snapshot(page(items, total, args.limit, args.offset))
async def results(self, args):
query = select(BacktestItem, BacktestResult).outerjoin(
BacktestResult, BacktestResult.item_id == BacktestItem.id
).where(BacktestItem.run_id == args.run_id)
if args.item_ids:
query = query.where(BacktestItem.id.in_(args.item_ids))
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = (await self.db.execute(query.order_by(BacktestItem.ordinal).limit(args.limit).offset(args.offset))).all()
return {"backtest_run_id": args.run_id, **page([item_summary(i, r) for i, r in rows], total, args.limit, args.offset)}
async def artifact(self, args):
from .service import ResearchError
item = await self.db.get(BacktestItem, args.item_id)
if not item:
raise ResearchError("NOT_FOUND", "候选不存在")
result = await self.db.get(BacktestResult, item.id)
if args.kind == "snapshot":
# Top-level entries retain complete nested values; no hidden string/list truncation.
entries = [{"key": k, "value": v} for k, v in sanitize(result.snapshot).items()] if result else []
return encode_snapshot({"item_id": item.id, "kind": args.kind,
"status": "available" if result else "not_available",
"observed_at": result.observed_at if result else None,
"complete": result.complete if result else False,
**page(entries[args.offset:args.offset + args.limit], len(entries), args.limit, args.offset)})
pnl = await self.db.get(Pnl, item.alpha_id) if item.alpha_id else None
points = pnl.points if pnl else []
points = [p for p in points if
(not args.date_from or p["date"][:10] >= args.date_from.isoformat()) and
(not args.date_to or p["date"][:10] <= args.date_to.isoformat())]
return encode_snapshot({"item_id": item.id, "alpha_id": item.alpha_id, "kind": args.kind,
"status": "available" if pnl else "not_cached", "fetched_at": pnl.fetched_at if pnl else None,
"units": "供应商原始累计值;未提供货币或规模单位",
**page(points[args.offset:args.offset + args.limit], len(points), args.limit, args.offset)})
+222
View File
@@ -0,0 +1,222 @@
"""Direct research operations; caller owns authorization, transaction and wake-up.
The account row serializes mutations with existing HTTP starts. Request records,
previews, runs and control events commit together; failed validation consumes no key.
"""
from collections import Counter
from uuid import uuid4
from sqlalchemy import func, select
from ..backtests.contracts import ControlInput, DraftInput, PreviewInput, Source, StartInput, fingerprint
from ..backtests.service import Backtests
from ..business import Business
from ..catalog.contracts import CatalogJobInput
from ..catalog.platform import platform_options, validate_platform_scope
from ..catalog.research_metadata import ResearchMetadata, availability_key
from ..catalog.service import Catalog
from ..models import Account, Alpha, BacktestItem, Job, JobItem, ResearchRequest, SimulationAttempt, now
from ..research.serialization import encode_snapshot
from ..research.workspace_contracts import FieldAvailabilityInput
from ..schemas import JobInput
from .contracts import DirectCandidate, History
from .queries import EvidenceQueries, page
class ResearchError(Exception):
def __init__(self, code, message, *, retryable=False, retry_after=None, affected_items=None):
super().__init__(message)
self.data = {"code": code, "message": message, "retryable": retryable,
"retry_after": retry_after, "affected_items": affected_items or []}
class ResearchAccess:
def __init__(self, db, principal, client, public_origin):
self.db, self.principal, self.client = db, principal, client
self.public_origin = public_origin.rstrip("/")
self.backtests = Backtests(db)
self.business = Business(db)
self.evidence = EvidenceQueries(db)
self.wake = None
def run_url(self, run_id):
return f"{self.public_origin}/#backtests?run_id={run_id}"
async def capabilities(self, args):
return {**await self.backtests.capabilities(), "max_candidates": 100,
"settings_schema": DirectCandidate.model_json_schema(),
"confirmation": "调用者须已获本批执行授权;直接提交后返回稳定运行 ID",
"duplicate_policies": ["reject", "rerun"], "permissions": sorted(self.principal.scopes),
"metadata_only": True, "actual_platform_allowance": None}
async def catalog(self, args):
data = await Catalog(self.db).search(args.filters, args.dataset_id)
return {**data, "scope": args.filters.model_dump(include={"region", "universe", "delay", "instrument_type"}),
"status": "available" if data["collection_version"] else "not_cached",
"has_more": data["offset"] + len(data["items"]) < data["total"]}
async def metadata(self, args):
q = args.query
metadata = ResearchMetadata(self.db)
if q.kind == "scopes":
return {"source": "worldquant_platform", **await platform_options(self.client)}
if q.kind == "operators":
data = await metadata.operators(q.q, q.category, limit=q.limit, offset=q.offset)
return {**data, "status": "available" if data["fetched_at"] else "not_cached",
"has_more": q.offset + len(data["items"]) < data["total"]}
if q.kind == "settings":
data = await metadata.get("settings")
items = data["content"].get("items", [])
return {"status": "available" if data["fetched_at"] else "not_cached",
"fetched_at": data["fetched_at"], **page(items[q.offset:q.offset+q.limit], len(items), q.limit, q.offset)}
data = await metadata.get(availability_key(q.field_id, q.scope))
return {**data, "status": data["content"].get("status", "unknown")}
async def refresh(self, args):
q = args.query
metadata = ResearchMetadata(self.db, self.client)
if q.kind == "catalog":
await validate_platform_scope(self.client, q.scope)
job = await Catalog(self.db).create_job(CatalogJobInput(scope=q.scope, dataset_id=q.dataset_id))
self.wake = "jobs"
return {"job_id": job.id, "status": job.status}
if q.kind == "pnl":
ids = sorted(set(q.alpha_ids))
existing = set(await self.db.scalars(select(Alpha.id).where(Alpha.id.in_(ids))))
if existing != set(ids):
raise ResearchError("NOT_FOUND", "部分 Alpha 尚未同步", affected_items=sorted(set(ids)-existing))
job = await self.business.create_sync_job(JobInput(kind="pnl_refresh", alpha_ids=ids))
self.wake = "jobs"
return {"job_id": job["id"], "status": job["status"]}
if q.kind == "operators":
data = await metadata.refresh_operators()
elif q.kind == "settings":
data = await metadata.refresh_settings()
else:
data = await metadata.refresh_availability(FieldAvailabilityInput(field_id=q.field_id, scope=q.scope))
# Refresh acknowledgment is bounded; complete content is available through paged reads.
return {"status": "completed", "key": data["key"], "fetched_at": data["fetched_at"],
"read_with": "get_research_metadata"}
async def refresh_job(self, args):
job = await self.db.get(Job, args.job_id)
if not job or job.kind not in {"catalog_sync", "field_sync", "pnl_refresh"}:
raise ResearchError("NOT_FOUND", "研究刷新任务不存在")
result = await self.business.get_job_status(args.job_id)
query = select(JobItem).where(JobItem.job_id == job.id, JobItem.error.is_not(None))
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
errors = list(await self.db.scalars(query.order_by(JobItem.alpha_id).limit(args.limit).offset(args.offset)))
result.pop("errors", None)
return {**result, "job_id": job.id, "artifact_reference": job.payload,
"errors": page([{"alpha_id": e.alpha_id, "error": e.error} for e in errors], total, args.limit, args.offset)}
async def history(self, args):
return await self.evidence.history(args)
async def previous(self, operation, args):
# PostgreSQL row lock is shared with HTTP start and catalog/job creation.
account = await self.db.scalar(select(Account).where(Account.id == self.principal.account_id).with_for_update())
if not account or account.wq_user_id != self.principal.wq_user_id:
raise ResearchError("ACCOUNT_MISMATCH", "平台账户绑定已变化")
digest = fingerprint(args.model_dump(mode="json", exclude={"idempotency_key"}))
row = await self.db.scalar(select(ResearchRequest).where(
ResearchRequest.account_id == account.id, ResearchRequest.operation == operation,
ResearchRequest.idempotency_key == args.idempotency_key))
if row and row.digest != digest:
raise ResearchError("IDEMPOTENCY_CONFLICT", "幂等键已用于不同内容")
return row, digest
async def remember(self, operation, args, digest, result):
result["_meta"] = {"schema_version": 1, "observed_at": now().isoformat(), "source": "system"}
self.db.add(ResearchRequest(id=str(uuid4()), account_id=self.principal.account_id,
operation=operation, idempotency_key=args.idempotency_key, digest=digest,
business_id=result["backtest_run_id"], response=encode_snapshot(result)))
await self.db.flush()
self.wake = "backtests"
return result
async def validate_settings(self, candidates):
snapshot = await ResearchMetadata(self.db).get("settings")
options = snapshot["content"].get("items", [])
if not snapshot["fetched_at"] or not options:
return {"settings_validation": "unknown", "reason": "设置快照未缓存;未验证平台组合", "field_validation": "unknown"}
invalid = []
for c in candidates:
s = c.settings
matches = [r for r in options if all(r.get(k) == v for k, v in {
"instrument_type": s.instrumentType, "region": s.region, "universe": s.universe, "delay": s.delay}.items())]
if not matches or all(r.get("neutralizations") and s.neutralization not in r["neutralizations"] for r in matches):
invalid.append(c.client_item_id)
if invalid:
raise ResearchError("UNSUPPORTED_SETTINGS", "已缓存平台设置不支持这些组合;可显式刷新后重试", affected_items=invalid)
return {"settings_validation": "cached", "fetched_at": snapshot["fetched_at"], "field_validation": "unknown"}
async def submit(self, args):
previous, digest = await self.previous("submit_backtests", args)
if previous:
return previous.response
seen, within = {}, []
for c in args.candidates:
h = fingerprint(c.platform_input())
if h in seen:
within.append({"client_item_id": c.client_item_id, "duplicate_of": seen[h]})
seen[h] = c.client_item_id
history = await self.evidence.history(History(candidates=args.candidates, limit=100))
if args.duplicate_policy == "reject" and (within or history["total"]):
raise ResearchError("DUPLICATE_INPUT", "发现完整输入重复;未创建运行。重跑须明确 duplicate_policy=rerun",
affected_items={"within_batch": within, "history": history, "read_with": "search_backtests"})
validation = await self.validate_settings(args.candidates)
if args.source.parent_run_id:
await self.backtests.run(args.source.parent_run_id)
source = Source(kind="mcp", **args.source.model_dump())
provenance = {"mcp_token_id": self.principal.token_id, "admin_id": self.principal.admin_id}
# preserve_source prevents the Chatbox-specific generating context rewriting MCP provenance.
backtests = Backtests(self.db, provenance)
preview = await backtests.preview(PreviewInput(inline=DraftInput(
name=args.name, source=source, candidates=args.candidates)), preserve_source=True)
result = await backtests.start(StartInput(preview_id=preview["preview_id"],
idempotency_key="mcp-" + str(uuid4())))
result = {**result, "input_digest": digest, "batch_count": preview["batch_count"],
"duplicates": {"within_batch": within, "historical_matches": history["total"]},
"validation": validation, "web_url": self.run_url(result["backtest_run_id"])}
return await self.remember("submit_backtests", args, digest, result)
async def run(self, args):
result = await self.backtests.run(args.run_id)
attempts = list(await self.db.scalars(select(SimulationAttempt).where(SimulationAttempt.run_id == args.run_id)))
result["submission_counts"] = {
"candidates": result["total"], "attempts": len(attempts),
"post_requests": sum(a.submit_count for a in attempts),
"confirmed_accepted_candidates": sum(len(a.payload) for a in attempts if a.progress_url),
"unknown_acceptance_candidates": sum(len(a.payload) for a in attempts if a.error_code == "submission_unknown" or (a.state == "submitting" and not a.progress_url)),
"actual_platform_consumption": None,
}
if args.after is not None:
result["events"] = await self.backtests.events(args.run_id, args.after, args.event_limit)
return {**result, "web_url": self.run_url(args.run_id)}
async def results(self, args):
await self.backtests.run(args.run_id)
if args.item_ids:
found = set(await self.db.scalars(select(BacktestItem.id).where(
BacktestItem.run_id == args.run_id, BacktestItem.id.in_(args.item_ids))))
if found != set(args.item_ids):
raise ResearchError("NOT_FOUND", "部分候选不属于此运行")
return await self.evidence.results(args)
async def artifact(self, args):
return await self.evidence.artifact(args)
async def control(self, args):
previous, digest = await self.previous("control_backtest", args)
if previous:
return previous.response
before = await self.backtests.run(args.run_id)
states = list(await self.db.scalars(select(SimulationAttempt.state).where(SimulationAttempt.run_id == args.run_id)))
result = await self.backtests.control(args.run_id, ControlInput(action=args.action, version=args.expected_version))
result["impact"] = {"remote_cancelled": False, "attempts_before": dict(Counter(states)),
"indefinite_account_block_cleared": args.action == "resume" and bool(before["scheduler"]["blocked_reason"])
and before["scheduler"]["blocked_until"] is None,
"note": "暂停/停止仅阻止后续提交;已提交模拟继续采集。recover 不重新提交。"}
return await self.remember("control_backtest", args, digest, result)