"""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 ..correlation import MIN_SAMPLES, THRESHOLD, WINDOW_YEARS 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, "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 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"}: 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 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)