feat: add MCP research access and browser key management
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Direct research interface shared by trusted application adapters."""
|
||||
@@ -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
|
||||
@@ -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)})
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user