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

362 lines
22 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 ..alphas import check_summary
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 ..correlation import MIN_SAMPLES, THRESHOLD, WINDOW_YEARS
from ..jobs import AUTH_KINDS
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 ..submission import CheckInput, correlation_allows_check, create_check_job, local_alpha, source
from ..submission import fingerprint as submission_fingerprint
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,
"templates": {
"create_with": "create_research_template", "required_scope": "research:write",
"authored_by": "caller", "max_source_items": 20,
"source_items_with": "get_backtest_results", "starts_backtests": False,
"web_url": f"{self.public_origin}/#templates",
},
"submission_check": {
"check_with": "check_submission", "read_with": "get_submission_check",
"job_with": "get_refresh_job", "max_targets": 1,
"description_writeback": True, "production_submission": False,
},
"self_correlation": {
"check_with": "check_self_correlation", "read_with": "get_self_correlation",
"job_with": "get_refresh_job", "max_targets": 100, "source": "local",
"reference_scope": "本地已同步的同地区已提交 Alpha,排除自身",
"method": "累计 PnL 日变化的 Pearson 相关系数,取带符号最大值",
"threshold": THRESHOLD, "min_samples": MIN_SAMPLES, "window_years": WINDOW_YEARS,
"platform_check": False,
}}
async def connection(self, args):
"""Read bounded connection evidence; never return credentials or session cookies."""
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", "平台账户绑定已变化")
job = None
if args.job_id:
job = await self.db.get(Job, args.job_id)
if not job or job.kind not in AUTH_KINDS:
raise ResearchError("NOT_FOUND", "认证任务不存在")
return {"connection_status": account.connection_status,
"session_authenticated": self.client.authenticated,
"credentials_configured": bool(account.email and account.password_encrypted),
"requires_human_verification": account.connection_status == "verification_required",
"web_url": f"{self.public_origin}/",
"job": {"job_id": job.id, "kind": job.kind, "status": job.status} if job else None,
"next_step": "需要人工验证时在网页账户连接面板完成,再调用 authenticate_worldquant(action=verify)"}
async def authenticate(self, args):
"""Queue authentication using saved credentials, sharing the HTTP account lock.
The caller commits the job and audit together before waking the runner.
Human challenges remain in the browser; no password or challenge URL is exposed.
"""
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", "平台账户绑定已变化")
if not account.email or not account.password_encrypted:
raise ResearchError("CREDENTIALS_NOT_CONFIGURED", "请先在网页保存 WorldQuant 账户配置")
job = await self.db.scalar(select(Job).where(
Job.kind.in_(AUTH_KINDS), Job.status.in_(("running", "queued"))))
if not job:
# Preserve a pending challenge instead of replacing its server-side session.
if args.action == "connect" and account.connection_status == "verification_required":
raise ResearchError("VERIFICATION_REQUIRED", "请在网页完成人工验证,再调用 action=verify")
job = Job(id=str(uuid4()), kind=args.action, payload={}, checkpoint={})
self.db.add(job)
if args.action == "connect":
account.connection_status, account.connection_error = "connecting", None
await self.db.flush()
self.wake = "jobs"
return {"job_id": job.id, "status": job.status, "action": job.kind,
"read_with": "get_worldquant_connection", "web_url": f"{self.public_origin}/"}
async def preparations(self, args):
from ..preparations.service import Preparations
return await Preparations(self.db).list(args.q, args.scope_key, args.limit, args.offset)
async def preparation(self, args):
from ..preparations.service import Preparations
service = Preparations(self.db)
row = await service.get(args.id, args.version, lock=True)
return {"collection": await service.output(row),
"fields": await service.members(row.id, args.q, None, args.limit, args.offset)}
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", "self_correlation", "submission_check"}:
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 check_self_correlation(self, args):
"""Queue local checks for synced IDs; caller commits before waking the runner.
The shared job service deduplicates active batches. Missing PnL is fetched
by the durable runner, so slow upstream reads do not hold the MCP call.
"""
ids = sorted(set(args.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="self_correlation", alpha_ids=ids))
self.wake = "jobs"
return {"job_id": job["id"], "status": job["status"], "alpha_ids": ids,
"source": "local", "job_with": "get_refresh_job", "read_with": "get_self_correlation"}
async def self_correlation(self, args):
"""Read the latest local result without fetching PnL or starting a check."""
data = await self.business.get_self_correlation(args.alpha_id)
status = "not_cached" if not data["cached"] else "stale" if data["result"]["stale"] else "available"
return {"alpha_id": args.alpha_id, "source": "local", "status": status, **data}
async def submission_check_context(self, args):
"""Read cached review context and check evidence without upstream requests."""
alpha = await local_alpha(self.db, args.alpha_id)
context = source(alpha.raw)
job = await self.db.scalar(select(Job).where(
Job.kind == "submission_check", Job.payload["alpha_ids"][0].as_string() == args.alpha_id
).order_by(Job.created_at.desc()).limit(1))
return {"alpha_id": args.alpha_id, "snapshot": submission_fingerprint(context),
**context, "descriptions": {key: item["description"] for key, item in context["sections"].items()},
"can_check": alpha.status == "UNSUBMITTED" and await correlation_allows_check(self.db, args.alpha_id),
"checks": alpha.checks, "check_summary": check_summary(alpha.checks, check_type=alpha.check_type), "source": "local_cache", "production_submission": False,
"job_id": job.id if job else None, "job_status": job.status if job else None,
"checked_at": job.checkpoint.get("checked_at") if job else None}
async def check_submission(self, args):
"""Queue the shared reviewed-description/check flow; never submit an Alpha."""
job = await create_check_job(self.db, args.alpha_id, CheckInput(
snapshot=args.snapshot, descriptions=args.descriptions))
self.wake = "jobs"
return {"job_id": job.id, "status": job.status, "alpha_id": args.alpha_id,
"production_submission": False, "job_with": "get_refresh_job",
"read_with": "get_submission_check"}
async def history(self, args):
return await self.evidence.history(args)
async def create_template(self, args):
"""Persist the caller's template and evidence without model or queue execution."""
from .templates import create_template
previous, digest = await self.previous("create_research_template", args)
if previous:
return previous.response
result = await create_template(self.db, args, self.principal)
result["web_url"] = f"{self.public_origin}/#templates"
return await self.remember("create_research_template", args, digest, result,
business_id=result["template_id"])
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, *, business_id, wake=None):
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=business_id, response=encode_snapshot(result)))
await self.db.flush()
self.wake = wake
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,
preparation_refs=args.preparation_refs)), 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,
business_id=result["backtest_run_id"], wake="backtests")
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,
business_id=result["backtest_run_id"], wake="backtests")