382 lines
23 KiB
Python
382 lines
23 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 ..worldquant import WqError
|
|
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 pyramid_distribution(self, args):
|
|
"""Read the supplied date's full quarter using the three-Alpha completion rule."""
|
|
from .pyramids import distribution, quarter_period
|
|
|
|
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", "平台账户绑定已变化")
|
|
period = quarter_period(args.current_date)
|
|
try:
|
|
raw = await self.client.get_pyramid_alphas(period["start_date"], period["end_date"])
|
|
groups = distribution(raw, args.region, args.delay)
|
|
except WqError as exc:
|
|
raise ResearchError(exc.code.upper(), str(exc), retryable=True) from None
|
|
except ValueError as exc:
|
|
raise ResearchError("INVALID_PLATFORM_DATA", str(exc)) from None
|
|
return {"region": args.region, "delay": args.delay, "threshold": 3,
|
|
"current_date": args.current_date.isoformat(), "period": period,
|
|
"source": "worldquant_platform", **groups}
|
|
|
|
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")
|