Files
worldquant-alpha-system/backend/app/research_access/service.py
T

223 lines
13 KiB
Python

"""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)