refactor: unify data preparations and research input snapshots
Deploy production / deploy (push) Successful in 53s

This commit is contained in:
yuxuanhui
2026-09-12 01:24:02 +08:00
parent 849f86fef7
commit 394438e753
82 changed files with 4146 additions and 2076 deletions
+3 -1
View File
@@ -47,6 +47,8 @@ class PageContext(Contract):
"alphas",
"account",
"datasets",
"fields",
"preparations",
"backtests",
"operators",
"templates",
@@ -62,7 +64,7 @@ class PageContext(Contract):
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
field_id: str | None = Field(default=None, min_length=1, max_length=200)
collection_version: str | None = Field(default=None, min_length=1, max_length=36)
template_input_id: str | None = Field(default=None, min_length=1, max_length=36)
input_snapshot_id: str | None = Field(default=None, min_length=1, max_length=36)
unsaved_field_selection: bool = False
backtest_run_id: str | None = Field(default=None, max_length=36)
backtest_preview_id: str | None = Field(default=None, max_length=36)
+7 -1
View File
@@ -6,6 +6,7 @@ from typing import Literal
from pydantic import Field, field_validator, model_validator
from ..preparations.contracts import PreparationReference
from ..schemas import Contract
@@ -48,13 +49,18 @@ class Source(Contract):
kind: str = Field(default="manual", min_length=1, max_length=100)
reference: str | None = Field(default=None, max_length=200)
batch_id: str | None = Field(default=None, max_length=200)
template_input_id: str | None = Field(default=None, max_length=200)
input_snapshot_ids: list[str] = Field(default_factory=list, max_length=100)
input_snapshot_id: str | None = Field(default=None, max_length=200)
research_id: str | None = Field(default=None, max_length=200)
parent_run_id: str | None = Field(default=None, max_length=36)
hypothesis: str | None = Field(default=None, max_length=2000)
class DraftInput(Contract):
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
input_ids: list[str] = Field(default_factory=list, max_length=20)
name: str = Field(min_length=1, max_length=200)
source: Source = Field(default_factory=Source)
candidates: list[Candidate] = Field(min_length=1, max_length=10000)
+27 -1
View File
@@ -106,8 +106,33 @@ class Backtests:
"limits": "并发和批大小为本地配置,并非平台可用额度;账户外提交不在预算内",
}
async def bind_preparations(self, body):
from ..preparations.service import Preparations
from ..research.expressions import analyze
if not body.preparation_refs and not body.input_ids:
return
await Preparations(self.db).bind(body)
snapshots = [await Preparations(self.db).snapshot(i) for i in body.input_ids]
for candidate in body.candidates:
scope = dict(instrument_type=candidate.settings.instrumentType, region=candidate.settings.region,
universe=candidate.settings.universe, delay=candidate.settings.delay)
if any(s["scope"] != scope for s in snapshots):
raise HTTPException(422, "数据准备集合与回测范围不一致")
fields = {}
for snapshot in snapshots:
for field, kind in snapshot["field_types"].items():
if field in fields and fields[field] != kind:
raise HTTPException(422, "输入字段类型冲突")
fields[field] = kind
validation = analyze(candidate.expression, fields)
if validation["syntax"] or validation["types"]:
raise HTTPException(422, ";".join(validation["syntax"] + validation["types"]))
body.source.input_snapshot_ids = body.input_ids
body.source.input_snapshot_id = body.input_ids[0] if len(body.input_ids) == 1 else None
async def save_draft(self, body, draft_id=None):
data = body.model_dump(mode="json", exclude={"version"})
await self.bind_preparations(body)
data = body.model_dump(mode="json", exclude={"version", "preparation_refs", "input_ids"})
if draft_id:
changed = await self.db.execute(
update(BacktestDraft)
@@ -168,6 +193,7 @@ class Backtests:
producer. ai_context separately identifies whoever starts the execution.
"""
if body.inline:
await self.bind_preparations(body.inline)
data = body.inline.model_dump(mode="json")
if self.ai_context and not preserve_source:
data["source"] = {
+9 -1
View File
@@ -263,9 +263,17 @@ class Business:
return {"ok": True, "job_id": job_id}
async def retry_job(self, job_id):
job = await self.db.scalar(select(Job).where(Job.id == job_id).with_for_update())
# Match create_job's lock order so retry and a fresh scheduled run share one scope owner.
await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
job = await self.db.scalar(select(Job).where(Job.id == job_id).with_for_update()
.execution_options(populate_existing=True))
if not job:
raise HTTPException(404, "任务不存在")
if job.kind == "catalog_full_sync":
active = await self.db.scalars(select(Job).where(Job.kind == job.kind, Job.status.in_(ACTIVE)))
for existing in active:
if existing.payload == job.payload and (existing.id != job.id or job.status in ("queued", "running")):
return JobOutput.model_validate(existing).model_dump(mode="json")
if job.status not in (
"failed",
"cancelled",
+1 -27
View File
@@ -3,7 +3,7 @@
from datetime import datetime, timezone
from typing import Annotated, Literal
from pydantic import AfterValidator, BaseModel, Field, model_validator
from pydantic import AfterValidator, BaseModel, Field
from ..schemas import Contract
@@ -54,20 +54,6 @@ class NoteInput(Contract):
version: int = Field(ge=1)
class InputPreparation(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
collection_version: str
selection: Literal["all", "explicit"] = "all"
excluded_ids: list[str] = Field(default_factory=list, max_length=100000)
@model_validator(mode="after")
def valid_selection(self):
if self.selection == "all" and self.excluded_ids:
raise ValueError("全部字段不能同时提供排除项")
return self
class NoteOutput(BaseModel):
note: str
version: int
@@ -107,18 +93,6 @@ class CatalogPage(BaseModel):
field_types: list[str] = Field(default_factory=list)
class InputOutput(BaseModel):
id: str
status: Literal["draft"] = "draft"
scope: Scope
dataset_id: str
collection_version: str
selection: str
field_ids: list[str]
field_types: dict[str, str | None]
created_at: UTCTimestamp
class CollectionOutput(BaseModel):
collection_version: str | None
field_ids: list[str]
-20
View File
@@ -12,8 +12,6 @@ from .contracts import (
CatalogPage,
CollectionOutput,
EntryOutput,
InputOutput,
InputPreparation,
NoteInput,
NoteOutput,
Scope,
@@ -76,24 +74,6 @@ async def sync(request: Request, body: CatalogJobInput):
return result
@router.post("/inputs", status_code=201, response_model=InputOutput)
async def prepare(request: Request, body: InputPreparation):
async with request.app.state.sessions.begin() as db:
return await Catalog(db).prepare(body)
@router.get("/inputs", response_model=list[InputOutput])
async def inputs(request: Request, scope: Annotated[Scope, Query()]):
async with request.app.state.sessions() as db:
return await Catalog(db).inputs(scope)
@router.get("/inputs/{input_id}", response_model=InputOutput)
async def get_input(request: Request, input_id: str):
async with request.app.state.sessions() as db:
return await Catalog(db).input(input_id)
@router.get("/datasets/{dataset_id}/collection", response_model=CollectionOutput)
async def collection(request: Request, dataset_id: str, scope: Annotated[Scope, Query()]):
async with request.app.state.sessions() as db:
+5 -65
View File
@@ -4,7 +4,6 @@ The dataset row serializes collection publication and draft creation on PostgreS
No page filters participate in template input selection.
"""
from datetime import timezone
from uuid import uuid4
from fastapi import HTTPException
@@ -18,7 +17,6 @@ from ..models import (
CatalogNote,
CatalogScope,
Job,
TemplateInput,
now,
)
from ..schemas import JobOutput
@@ -164,13 +162,13 @@ class Catalog:
raise HTTPException(409, "研究备注已被修改;当前草稿已保留,请载入最新记录后重新保存")
return dict(note=body.note, version=body.version + 1, updated_at=now())
async def create_job(self, body):
async def create_job(self, body, *, full=False):
account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
if not account.password_encrypted or account.connection_status in ("disconnected", "error"):
raise HTTPException(409, "请先连接 WorldQuant")
if body.dataset_id:
await self.dataset(body.scope, body.dataset_id)
kind = "field_sync" if body.dataset_id else "catalog_sync"
kind = "catalog_full_sync" if full else "field_sync" if body.dataset_id else "catalog_sync"
payload = body.model_dump(mode="json")
jobs = (
await self.db.scalars(
@@ -190,7 +188,7 @@ class Catalog:
job = Job(id=str(uuid4()), kind=kind, payload=payload)
self.db.add(job)
await self.db.flush()
self.db.add(CatalogBatch(id=job.id, scope_key=body.scope.key(), dataset_id=body.dataset_id))
self.db.add(CatalogBatch(id=job.id, job_id=job.id, scope_key=body.scope.key(), dataset_id=body.dataset_id))
await self.db.flush()
return JobOutput.model_validate(job)
@@ -210,64 +208,6 @@ class Catalog:
)
return dict(collection_version=dataset.field_version, field_ids=ids)
async def prepare(self, body):
dataset = await self.dataset(body.scope, body.dataset_id, lock=True)
if not dataset.field_version or dataset.field_version != body.collection_version:
raise HTTPException(409, "字段集合未完成或版本已变化,请重新读取后准备输入")
batch = await self.db.get(CatalogBatch, dataset.field_version)
if not batch.complete or batch.scope_key != body.scope.key() or batch.dataset_id != body.dataset_id:
raise HTTPException(409, "字段集合不完整")
entries = (
await self.db.scalars(
select(CatalogEntry).where(CatalogEntry.batch_id == batch.id).order_by(CatalogEntry.id)
)
).all()
fields = {e.id: e.field_type for e in entries}
excluded = set(body.excluded_ids)
if excluded - fields.keys():
raise HTTPException(422, "排除项含未知、跨范围或其他数据集字段")
chosen = {key: value for key, value in fields.items() if key not in excluded}
if not chosen:
raise HTTPException(422, "模板输入至少需要一个字段")
row = TemplateInput(
id=str(uuid4()),
scope_key=body.scope.key(),
dataset_id=body.dataset_id,
collection_version=batch.id,
selection=body.selection,
field_ids=list(chosen),
field_types=chosen,
)
self.db.add(row)
await self.db.flush()
return await self.input(row.id)
async def input(self, input_id):
row = await self.db.get(TemplateInput, input_id)
if not row:
raise HTTPException(404, "输入草稿不存在")
scope = await self.db.get(CatalogScope, row.scope_key)
return dict(
id=row.id,
status="draft",
scope=scope.scope,
dataset_id=row.dataset_id,
collection_version=row.collection_version,
selection=row.selection,
field_ids=row.field_ids,
field_types=row.field_types,
created_at=row.created_at.replace(tzinfo=timezone.utc)
if row.created_at.tzinfo is None
else row.created_at,
)
async def inputs(self, scope):
ids = (
await self.db.scalars(
select(TemplateInput.id)
.where(TemplateInput.scope_key == scope.key())
.order_by(TemplateInput.created_at.desc())
.limit(100)
)
).all()
return [await self.input(i) for i in ids]
from ..preparations.service import Preparations
return await Preparations(self.db).snapshot(input_id)
+82 -14
View File
@@ -4,6 +4,7 @@ import asyncio
import math
import re
from urllib.parse import parse_qs, urlparse
from uuid import uuid4
from sqlalchemy import select
@@ -64,14 +65,15 @@ def normalize(raw, dataset_id):
)
async def sync_catalog(runner, job_id, payload):
async def sync_catalog(runner, job_id, payload, *, batch_id=None, full=False):
scope = Scope.model_validate(payload["scope"])
dataset_id = payload.get("dataset_id")
batch_id = batch_id or job_id
async with runner.sessions() as db:
checkpoint = (await db.get(Job, job_id)).checkpoint
if checkpoint.get("done"):
return
offset = checkpoint.get("offset", 0)
batch = await db.get(CatalogBatch, batch_id)
if batch.complete:
return
offset = batch.offset
while True:
await runner.checkpoint(job_id, {"next_retry_at": None})
raw = await runner.client.catalog_page(scope.model_dump(), dataset_id, offset)
@@ -97,12 +99,12 @@ async def sync_catalog(runner, job_id, payload):
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
batch = await db.get(CatalogBatch, job_id)
batch = await db.get(CatalogBatch, batch_id)
added = 0
for entry in entries:
if await db.get(CatalogEntry, (job_id, entry["id"])):
if await db.get(CatalogEntry, (batch_id, entry["id"])):
continue
db.add(CatalogEntry(batch_id=job_id, **entry))
db.add(CatalogEntry(batch_id=batch_id, **entry))
await db.flush()
added += 1
owner = dataset_id or entry["id"]
@@ -112,25 +114,28 @@ async def sync_catalog(runner, job_id, payload):
if rows and not added:
raise WqError("平台分页重复且未前进,已保留进度", "invalid_response")
batch.count += added
job.processed = batch.count
if not full:
job.processed = batch.count
offset += len(rows)
job.checkpoint = dict(offset=offset, done=not more)
batch.offset = offset
job.checkpoint = {**job.checkpoint, "offset": offset, "done": not more, "current_field_count": batch.count}
job.updated_at = now()
if not more:
batch.complete, batch.completed_at = True, now()
job.total = batch.count
if not full:
job.total = batch.count
if dataset_id:
dataset = await db.scalar(
select(CatalogDataset)
.where(CatalogDataset.scope_key == scope.key(), CatalogDataset.id == dataset_id)
.with_for_update()
)
dataset.field_version = job_id
dataset.field_version = batch_id
else:
scope_row = await db.get(CatalogScope, scope.key())
scope_row.catalog_version, scope_row.synced_at = job_id, now()
scope_row.catalog_version, scope_row.synced_at = batch_id, now()
ids = (
await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id))
await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == batch_id))
).all()
for item_id in ids:
if not await db.get(CatalogDataset, (scope.key(), item_id)):
@@ -138,3 +143,66 @@ async def sync_catalog(runner, job_id, payload):
await db.commit()
if not more:
return
async def sync_full_catalog(runner, job_id, payload):
"""Resume each dataset batch independently; only publish complete enumerations."""
scope = Scope.model_validate(payload["scope"])
options = await runner.client.get_platform_setting_options()
if not any(r["instrument_type"] == scope.instrument_type and r["region"] == scope.region
and r["delay"] == scope.delay and scope.universe in r["universes"]
for r in options["instrument_options"]):
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.checkpoint = {**job.checkpoint, "error_code": "invalid_scope"}
raise WqError("平台不支持该研究范围", "invalid_scope")
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.checkpoint = {**{k: v for k, v in job.checkpoint.items() if k != "error_code"}, "phase": "catalog"}
await sync_catalog(runner, job_id, {"scope": payload["scope"]}, full=True)
async with runner.sessions() as db:
ids = list(await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id)
.order_by(CatalogEntry.id)))
job = await db.get(Job, job_id)
failures = dict(job.checkpoint.get("failures", {}))
completed = 0
await runner.checkpoint(job_id, {"total": len(ids)})
for dataset_id in ids:
async with runner.sessions.begin() as db:
batch = await db.scalar(select(CatalogBatch).where(CatalogBatch.job_id == job_id,
CatalogBatch.dataset_id == dataset_id))
if not batch:
batch = CatalogBatch(id=str(uuid4()), job_id=job_id, scope_key=scope.key(), dataset_id=dataset_id)
db.add(batch)
await db.flush()
batch_id, complete = batch.id, batch.complete
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
job.checkpoint = {**job.checkpoint, "phase": "fields", "dataset_id": dataset_id,
"datasets_completed": completed, "datasets_total": len(ids),
"offset": batch.offset, "current_field_count": batch.count}
if not complete:
try:
await sync_catalog(runner, job_id, {"scope": payload["scope"], "dataset_id": dataset_id},
batch_id=batch_id, full=True)
except WqError as exc:
if exc.code in ("disconnected", "authentication_failed", "identity_mismatch", "verification_required", "network_error"):
raise
failures[dataset_id] = str(exc)
if complete or (await _batch_complete(runner, batch_id)):
failures.pop(dataset_id, None)
completed += 1
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.processed, job.failed = completed, len(failures)
job.checkpoint = {**job.checkpoint, "datasets_completed": completed, "failures": failures}
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.error = f"{len(failures)} 个数据集同步失败" if failures else None
job.checkpoint = {**job.checkpoint, "phase": "finished"}
async def _batch_complete(runner, batch_id):
async with runner.sessions() as db:
return (await db.get(CatalogBatch, batch_id)).complete
+69
View File
@@ -3,6 +3,7 @@
import argparse
import asyncio
import getpass
import math
from sqlalchemy import delete, update
@@ -59,6 +60,63 @@ async def token_command(args):
await engine.dispose()
async def catalog_sync_command(args):
"""Enqueue on the existing runner, then observe without owning the upstream session."""
import json
import time
from fastapi import HTTPException
from .business import Business
from .catalog.contracts import CatalogJobInput, Scope
from .catalog.service import Catalog
from .models import Job
engine, sessions = create_database(Settings().database_url)
try:
async with sessions.begin() as db:
if args.resume_job:
if args.region or args.universe or args.delay is not None:
raise ValueError("--resume-job 不能同时指定新范围")
job = await db.get(Job, args.resume_job)
if not job or job.kind != "catalog_full_sync":
raise ValueError("只能恢复已有全量目录任务")
result = await Business(db).retry_job(job.id)
job_id = result["id"]
else:
if not args.region or not args.universe or args.delay is None:
raise ValueError("需要 --region、--universe 和 --delay")
scope = Scope(instrument_type=args.instrument_type, region=args.region,
universe=args.universe, delay=args.delay)
job = await Catalog(db).create_job(CatalogJobInput(scope=scope), full=True)
job_id = job.id
deadline, previous = time.monotonic() + args.wait_timeout, None
while True:
async with sessions() as db:
job = await db.get(Job, job_id)
data = dict(job_id=job.id, status=job.status, processed=job.processed,
total=job.total, failed=job.failed, checkpoint=job.checkpoint, error=job.error)
current = json.dumps(data, ensure_ascii=False, sort_keys=True)
if current != previous:
print(current, flush=True)
previous = current
if job.status == "completed":
return 0
if job.status in ("failed", "completed_with_errors", "cancelled"):
return 2 if job.checkpoint.get("error_code") == "invalid_scope" else 1
if job.status in ("waiting_auth", "waiting_connection"):
return 3
if time.monotonic() >= deadline:
print(f"等待超时;后台任务 {job_id} 继续执行", flush=True)
return 4
await asyncio.sleep(min(5, max(0, deadline - time.monotonic())))
except HTTPException as exc:
print(str(exc.detail), flush=True)
return 3 if exc.status_code == 409 else 1
finally:
await engine.dispose()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
commands = parser.add_subparsers(dest="command", required=True)
@@ -70,8 +128,19 @@ if __name__ == "__main__":
commands.add_parser("mcp-token-list")
revoke = commands.add_parser("mcp-token-revoke")
revoke.add_argument("token_id")
sync = commands.add_parser("catalog-sync", help="全量同步一个范围的数据集及全部字段")
sync.add_argument("--region")
sync.add_argument("--universe")
sync.add_argument("--delay", type=int, choices=range(0, 10))
sync.add_argument("--instrument-type", default="EQUITY")
sync.add_argument("--resume-job")
sync.add_argument("--wait-timeout", type=float, default=21600)
args = parser.parse_args()
if args.command == "catalog-sync" and (not math.isfinite(args.wait_timeout) or args.wait_timeout <= 0):
parser.error("--wait-timeout 必须大于 0")
try:
if args.command == "catalog-sync":
raise SystemExit(asyncio.run(catalog_sync_command(args)))
asyncio.run(reset_password() if args.command == "reset-password" else token_command(args))
except ValueError as exc:
parser.error(str(exc))
+5 -1
View File
@@ -255,6 +255,10 @@ class Runner:
await self.ensure_connected(force=kind == "connect")
if kind in ("connect", "profile"):
await self.refresh_profile()
elif kind == "catalog_full_sync":
from .catalog.sync import sync_full_catalog
await sync_full_catalog(self, job_id, payload)
elif kind in ("catalog_sync", "field_sync"):
from .catalog.sync import sync_catalog
@@ -298,7 +302,7 @@ class Runner:
await self.checkpoint(
job_id,
{
"status": "waiting_connection" if waiting else "failed",
"status": "waiting_connection" if waiting or (kind == "catalog_full_sync" and exc.code == "network_error") else "failed",
"error": str(exc),
"next_retry_at": None,
},
+2
View File
@@ -27,6 +27,7 @@ from .db import create_database
from .jobs import AUTH_KINDS, Runner, create_job
from .mcp_api.token_routes import router as mcp_token_router
from .models import Account, Admin, BacktestConfig, Job, JobItem, LoginSession
from .preparations.routes import router as preparations_router
from .research.routes import router as research_router
from .research.runtime import ResearchRuntime
from .schemas import (
@@ -484,6 +485,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
app.include_router(backtest_router)
app.include_router(api)
app.include_router(catalog_router)
app.include_router(preparations_router)
app.include_router(research_catalog_router)
app.include_router(research_router)
app.include_router(ai_router(ai_runtime))
+3 -1
View File
@@ -21,6 +21,8 @@ from ..research_access.service import ResearchAccess, ResearchError
# Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
TOOLS = {
"search_data_preparations": (c.PreparationSearch, "preparations", "research:read", "分页查询数据准备集合,返回固定范围、字段数及版本。研究可使用多个集合,各集合范围独立。"),
"get_data_preparation": (c.PreparationRead, "preparation", "research:read", "按集合 ID 与版本分页预览字段、类型、描述和数据集归属。提交回测时携带 preparation_refs,由服务端核对版本并固定独立输入快照;空集合不能用于研究。"),
"create_research_template": (c.CreateTemplate, "create_template", "research:write", "将调用方大模型研究后自行总结的参数化模板保存到模板工坊,供用户后续批量回测。先用 get_backtest_results 阅读实际指标和检查,选择 1–20 个已完成采集的 source_item_ids,并说明 hypothesis;不要把 completed 当作检查通过。template 使用 {name} 占位符及逐一对应的 variables,字段变量须声明 MATRIX/VECTOR/GROUP,VECTOR 聚合须明确写入表达式。提供唯一名称和 idempotency_key,可附 reference。返回模板 ID、版本和理论组合数;仅核验结构及来源,不验证所有参数组合,不再次调用模型、不执行回测、不覆盖已有模板。"),
"get_submission_check": (c.SelfCorrelationReference, "submission_check_context", "research:read", "读取已导入 Alpha 的表达式、Description、snapshot 和缓存检查结果;不发起检查。先核对或生成三段 Description,再调用 check_submission。"),
"check_submission": (c.SubmissionCheck, "check_submission", "research:refresh", "对单个待提交 Alpha 写回已获用户授权的 Description 并调用平台 GET /check,返回 job_id。须先用 get_submission_check 获取 snapshot;保留本地自相关门槛和冲突保护。通过 get_refresh_job 查进度、get_submission_check 读结果。无论检查结果如何,都不会调用 /submit 或正式提交 Alpha。"),
@@ -34,7 +36,7 @@ TOOLS = {
"check_self_correlation": (c.SelfCorrelationCheck, "check_self_correlation", "research:refresh", "对 1–100 个已导入 Alpha 发起本地自相关检查,返回 job_id。与本地同地区已提交 Alpha 比较,排除自身;优先用缓存,缺失 PnL 自动补取。需先同步已提交 Alpha;不调用平台提交检查。用 get_refresh_job 查进度、get_self_correlation 读结果。"),
"get_self_correlation": (c.SelfCorrelationReference, "self_correlation", "research:read", "读取指定 Alpha 最新的本地自相关缓存,包括最大相关系数、样本覆盖与 stale 状态;无缓存不自动检查,结果不等同于平台提交资格。"),
"search_backtests": (c.History, "history", "research:read", "分页查历史候选与固定设置;candidates 按完整输入精确匹配,不推断数学等价。"),
"submit_backtests": (c.Submit, "submit", "backtests:execute", "执行用户已授权的固定批次,自动留痕并立即返回运行 ID。每项必须完整设置;重复默认拒绝,rerun 明确重跑。不需要研究资产。"),
"submit_backtests": (c.Submit, "submit", "backtests:execute", "执行用户已授权的固定批次,自动留痕并立即返回运行 ID。可携带 preparation_refs 选择集合,版本变化须重新读取;每项必须完整设置;重复默认拒绝,rerun 明确重跑。不需要研究资产。"),
"get_backtest": (c.RunReference, "run", "research:read", "读取真实运行进度、提交数量和可选增量事件;受理不等于成功。"),
"get_backtest_results": (c.Results, "results", "research:read", "分页读取固定快照指标、全部非通过检查及三层状态;缺失指标不补零。"),
"get_backtest_artifact": (c.Artifact, "artifact", "research:read", "分页读取候选脱敏快照的顶层键值或独立采集的 PnL;缺缓存不自动刷新。"),
+31 -9
View File
@@ -340,7 +340,9 @@ class CatalogScope(Base):
class CatalogBatch(Base):
__tablename__ = "catalog_batches"
id: Mapped[str] = mapped_column(ForeignKey("sync_jobs.id"), primary_key=True)
id: Mapped[str] = mapped_column(String(36), primary_key=True)
job_id: Mapped[str] = mapped_column(ForeignKey("sync_jobs.id"), index=True)
offset: Mapped[int] = mapped_column(Integer, default=0)
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True)
dataset_id: Mapped[str | None] = mapped_column(String(200))
complete: Mapped[bool] = mapped_column(Boolean, default=False)
@@ -386,16 +388,36 @@ class CatalogNote(Base):
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
class TemplateInput(Base):
__tablename__ = "template_inputs"
class DataPreparation(Base):
"""Editable collection; scope never changes after creation."""
__tablename__ = "data_preparations"
id: Mapped[str] = mapped_column(String(36), primary_key=True)
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True)
dataset_id: Mapped[str] = mapped_column(String(200))
collection_version: Mapped[str] = mapped_column(ForeignKey("catalog_batches.id"))
selection: Mapped[str] = mapped_column(String(20))
field_ids: Mapped[list] = mapped_column(JSON)
field_types: Mapped[dict] = mapped_column(JSON)
name: Mapped[str] = mapped_column(String(200))
note: Mapped[str] = mapped_column(Text, default="")
scope_key: Mapped[str] = mapped_column(String(200), index=True)
scope: Mapped[dict] = mapped_column(JSON)
version: Mapped[int] = mapped_column(Integer, default=1)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
class PreparationField(Base):
__tablename__ = "preparation_fields"
preparation_id: Mapped[str] = mapped_column(ForeignKey("data_preparations.id", ondelete="CASCADE"), primary_key=True)
field_id: Mapped[str] = mapped_column(String(200), primary_key=True)
dataset_id: Mapped[str] = mapped_column(String(200), index=True)
content: Mapped[dict] = mapped_column(JSON)
class ResearchInputSnapshot(Base):
"""Self-contained research input: deletion of its preparation cannot invalidate it."""
__tablename__ = "research_input_snapshots"
id: Mapped[str] = mapped_column(String(36), primary_key=True)
preparation_id: Mapped[str] = mapped_column(String(36), index=True)
preparation_version: Mapped[int] = mapped_column(Integer)
content: Mapped[dict] = mapped_column(JSON)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
__table_args__ = (UniqueConstraint("preparation_id", "preparation_version"),)
class CatalogResource(Base):
+1
View File
@@ -0,0 +1 @@
"""Data preparation collections and immutable research inputs."""
+80
View File
@@ -0,0 +1,80 @@
"""Shared collection and field-query contracts."""
from datetime import datetime
from typing import Literal
from pydantic import Field, model_validator
from ..catalog.contracts import Scope
from ..schemas import Contract
class FieldFilters(Scope):
q: str = Field(default="", max_length=300)
dataset_id: str | None = None
category: str | None = None
subcategory: str | None = None
field_type: str | None = None
coverage_min: float | None = Field(default=None, ge=0, le=1)
coverage_max: float | None = Field(default=None, ge=0, le=1)
user_count_min: int | None = Field(default=None, ge=0)
user_count_max: int | None = Field(default=None, ge=0)
alpha_count_min: int | None = Field(default=None, ge=0)
alpha_count_max: int | None = Field(default=None, ge=0)
synced_from: datetime | None = None
synced_to: datetime | None = None
sort: Literal["id", "name", "dataset_id", "coverage", "user_count", "alpha_count", "synced_at"] = "id"
direction: Literal["asc", "desc"] = "asc"
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
@model_validator(mode="after")
def ranges(self):
for key in ("coverage", "user_count", "alpha_count"):
low, high = getattr(self, key + "_min"), getattr(self, key + "_max")
if low is not None and high is not None and low > high:
raise ValueError("筛选下限不能超过上限")
return self
class FieldReference(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
field_id: str = Field(min_length=1, max_length=200)
source: Literal["local", "worldquant"] = "local"
collection_version: str | None = None
class PreparationCreate(Contract):
name: str = Field(min_length=1, max_length=200)
note: str = Field(default="", max_length=20000)
scope: Scope
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
class PreparationVersion(Contract):
version: int = Field(ge=1)
class PreparationEdit(PreparationVersion):
name: str = Field(min_length=1, max_length=200)
note: str = Field(default="", max_length=20000)
class MemberChange(PreparationVersion):
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
remove_ids: list[str] = Field(default_factory=list, max_length=10000)
class PreparationReference(PreparationVersion):
id: str = Field(min_length=1, max_length=36)
class PreparationReferences(Contract):
items: list[PreparationReference] = Field(min_length=1, max_length=100)
class DatasetCopy(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
collection_version: str
+195
View File
@@ -0,0 +1,195 @@
"""Authenticated preparation and field-directory endpoints."""
from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import delete, select
from ..catalog.contracts import Scope
from ..models import CatalogScope, PreparationField, now
from ..security import require_auth
from .contracts import (
DatasetCopy,
FieldFilters,
MemberChange,
PreparationCreate,
PreparationEdit,
PreparationReferences,
PreparationVersion,
)
from .service import Preparations
router = APIRouter(prefix="/api/v1", dependencies=[Depends(require_auth)], tags=["data-preparations"])
@router.get("/catalog/local-scopes")
async def local_scopes(request: Request):
async with request.app.state.sessions() as db:
rows = await db.scalars(select(CatalogScope))
options = {}
for row in rows:
key = (row.scope["instrument_type"], row.scope["region"], row.scope["delay"])
option = options.setdefault(
key, {k: row.scope[k] for k in ("instrument_type", "region", "delay")}
)
option.setdefault("universes", []).append(row.scope["universe"])
return {"instrument_options": list(options.values())}
@router.get("/catalog/fields")
async def local_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
async with request.app.state.sessions() as db:
return await Preparations(db).fields(filters)
@router.get("/catalog/worldquant/fields")
async def online_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
async with request.app.state.sessions() as db:
return await Preparations(db, request.app.state.runner.client).online_fields(filters)
@router.get("/data-preparations")
async def preparations(
request: Request,
q: str = "",
scope_key: str | None = None,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
):
async with request.app.state.sessions() as db:
return await Preparations(db).list(q, scope_key, limit, offset)
@router.post("/data-preparations", status_code=201)
async def create(request: Request, body: PreparationCreate):
async with request.app.state.sessions.begin() as db:
service = Preparations(db, request.app.state.runner.client)
fields = await service.resolve_fields(body.scope, body.fields)
return await service.create(body.name, body.note, body.scope, fields)
@router.post("/data-preparations/from-dataset", status_code=201)
async def from_dataset(request: Request, body: DatasetCopy):
async with request.app.state.sessions.begin() as db:
return await Preparations(db).copy_dataset(body)
@router.post("/data-preparations/batch-delete")
async def batch_delete(request: Request, body: PreparationReferences):
async with request.app.state.sessions.begin() as db:
return await Preparations(db).remove(body.items)
@router.post("/data-preparations/freeze", status_code=201)
async def freeze(request: Request, body: PreparationReferences):
async with request.app.state.sessions.begin() as db:
return {"items": await Preparations(db).freeze(body.items)}
@router.get("/research/input-snapshots/{snapshot_id}")
async def snapshot(request: Request, snapshot_id: str):
async with request.app.state.sessions() as db:
return await Preparations(db).snapshot(snapshot_id)
@router.get("/data-preparations/{preparation_id}")
async def detail(request: Request, preparation_id: str):
async with request.app.state.sessions() as db:
service = Preparations(db)
return await service.output(await service.get(preparation_id))
@router.patch("/data-preparations/{preparation_id}")
async def edit(request: Request, preparation_id: str, body: PreparationEdit):
async with request.app.state.sessions.begin() as db:
service = Preparations(db)
row = await service.get(preparation_id, body.version, lock=True)
row.name, row.note, row.updated_at, row.version = body.name, body.note, now(), row.version + 1
return await service.output(row)
@router.delete("/data-preparations/{preparation_id}")
async def remove(request: Request, preparation_id: str, version: int = Query(ge=1)):
from .contracts import PreparationReference
async with request.app.state.sessions.begin() as db:
return await Preparations(db).remove([PreparationReference(id=preparation_id, version=version)])
@router.post("/data-preparations/{preparation_id}/copy", status_code=201)
async def copy(request: Request, preparation_id: str, body: PreparationVersion):
async with request.app.state.sessions.begin() as db:
service = Preparations(db)
row = await service.get(preparation_id, body.version, lock=True)
fields = [
f.content
for f in await db.scalars(
select(PreparationField).where(PreparationField.preparation_id == row.id)
)
]
return await service.create(
(row.name + " 副本")[:200], row.note, Scope.model_validate(row.scope), fields
)
@router.get("/data-preparations/{preparation_id}/fields")
async def members(
request: Request,
preparation_id: str,
q: str = "",
dataset_id: str | None = None,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
):
async with request.app.state.sessions() as db:
return await Preparations(db).members(preparation_id, q, dataset_id, limit, offset)
@router.patch("/data-preparations/{preparation_id}/fields")
async def change_members(request: Request, preparation_id: str, body: MemberChange):
async with request.app.state.sessions.begin() as db:
service = Preparations(db, request.app.state.runner.client)
row = await service.get(preparation_id, body.version, lock=True)
fields = await service.resolve_fields(Scope.model_validate(row.scope), body.fields)
await service.add(row, fields)
if body.remove_ids:
present = set(
await db.scalars(
select(PreparationField.field_id).where(
PreparationField.preparation_id == row.id,
PreparationField.field_id.in_(body.remove_ids),
)
)
)
if set(body.remove_ids) - present:
raise HTTPException(422, "移除项含不属于该集合的字段")
await db.execute(
delete(PreparationField).where(
PreparationField.preparation_id == row.id, PreparationField.field_id.in_(body.remove_ids)
)
)
row.version, row.updated_at = row.version + 1, now()
return await service.output(row)
@router.get("/data-preparations/{preparation_id}/selection")
async def selection(request: Request, preparation_id: str, version: int = Query(ge=1)):
async with request.app.state.sessions() as db:
service = Preparations(db)
row = await service.get(preparation_id, version)
fields = [
f.content
for f in await db.scalars(
select(PreparationField)
.where(PreparationField.preparation_id == row.id)
.order_by(PreparationField.field_id)
)
]
return {
**await service.output(row),
"fields": fields,
"field_ids": [f["id"] for f in fields],
"field_types": {f["id"]: f["field_type"] for f in fields},
"dataset_ids": sorted({f["dataset_id"] for f in fields}),
"preparation_ref": {"id": row.id, "version": row.version},
}
+450
View File
@@ -0,0 +1,450 @@
"""Collection operations own validation; callers own transactions and authorization."""
from uuid import uuid4
from fastapi import HTTPException
from sqlalchemy import delete, func, or_, select
from sqlalchemy.orm import aliased
from ..catalog.contracts import EntryOutput, Scope
from ..catalog.research_metadata import upstream
from ..catalog.service import Catalog
from ..catalog.sync import identifier, label, normalize
from ..models import (
CatalogBatch,
CatalogDataset,
CatalogEntry,
CatalogScope,
DataPreparation,
PreparationField,
ResearchInputSnapshot,
now,
)
from ..research.serialization import encode_snapshot
from .contracts import FieldFilters
def page(items, total, limit, offset, **extra):
return dict(
items=items, total=total, limit=limit, offset=offset, has_more=offset + len(items) < total, **extra
)
def contains(value):
return "%" + value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
class Preparations:
def __init__(self, db, client=None):
self.db, self.client = db, client
async def fields(self, filters):
scope = await self.db.get(CatalogScope, filters.key())
owner = aliased(CatalogEntry)
query = (
select(
CatalogEntry,
CatalogBatch.dataset_id,
owner.name.label("dataset_name"),
owner.category,
owner.subcategory,
)
.join(CatalogBatch, CatalogEntry.batch_id == CatalogBatch.id)
.join(
CatalogDataset,
(CatalogDataset.field_version == CatalogBatch.id)
& (CatalogDataset.scope_key == filters.key()),
)
.outerjoin(
owner,
(owner.id == CatalogBatch.dataset_id)
& (owner.batch_id == (scope.catalog_version if scope else None)),
)
.where(CatalogBatch.complete.is_(True))
)
if filters.q:
query = query.where(
or_(
*[
column.ilike(contains(filters.q), escape="\\")
for column in (
CatalogEntry.id,
CatalogEntry.name,
CatalogEntry.description,
CatalogBatch.dataset_id,
owner.name,
)
]
)
)
for key in ("dataset_id", "field_type", "category", "subcategory"):
value = getattr(filters, key)
column = (
CatalogBatch.dataset_id
if key == "dataset_id"
else func.coalesce(getattr(CatalogEntry, key), getattr(owner, key))
if key in ("category", "subcategory")
else getattr(CatalogEntry, key)
)
if value:
query = query.where(column == value)
for key in ("coverage", "user_count", "alpha_count"):
low, high = getattr(filters, key + "_min"), getattr(filters, key + "_max")
if low is not None:
query = query.where(getattr(CatalogEntry, key) >= low)
if high is not None:
query = query.where(getattr(CatalogEntry, key) <= high)
if filters.synced_from:
query = query.where(CatalogEntry.synced_at >= filters.synced_from)
if filters.synced_to:
query = query.where(CatalogEntry.synced_at <= filters.synced_to)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
column = (
CatalogBatch.dataset_id if filters.sort == "dataset_id" else getattr(CatalogEntry, filters.sort)
)
query = query.order_by(
(column.desc() if filters.direction == "desc" else column.asc()).nulls_last(),
CatalogBatch.dataset_id,
CatalogEntry.id,
)
rows = (await self.db.execute(query.limit(filters.limit).offset(filters.offset))).all()
items = [
{
**self.local_field(row, dataset_id, name, filters),
"category": row.category or category,
"subcategory": row.subcategory or subcategory,
}
for row, dataset_id, name, category, subcategory in rows
]
return page(items, total, filters.limit, filters.offset)
@staticmethod
def local_field(row, dataset_id, name, scope):
return encode_snapshot(
dict(
**EntryOutput.model_validate(row, from_attributes=True).model_dump(
exclude={"scope", "dataset_id", "collection_version"}
),
field_id=row.id,
dataset_id=dataset_id,
dataset_name=name or dataset_id,
collection_version=row.batch_id,
scope=scope.model_dump(include=set(Scope.model_fields)),
source="local",
fetched_at=row.synced_at,
)
)
async def online_fields(self, filters):
if self.client is None:
raise HTTPException(409, "请先连接 WorldQuant")
params = dict(
instrumentType=filters.instrument_type,
region=filters.region,
universe=filters.universe,
delay=filters.delay,
limit=filters.limit,
offset=filters.offset,
)
for key, remote in (("q", "search"), ("dataset_id", "dataset.id"), ("field_type", "type")):
value = getattr(filters, key)
if value:
params[remote] = value
for key, remote in (
("coverage", "coverage"),
("user_count", "userCount"),
("alpha_count", "alphaCount"),
):
for suffix, op in (("min", ">"), ("max", "<")):
value = getattr(filters, key + "_" + suffix)
if value is not None:
params[remote + op] = value
raw = await upstream(self.client.get("/data-fields", params))
rows = raw.get("results")
if not isinstance(rows, list):
raise HTTPException(502, "平台字段列表格式无法识别")
items = []
for row in rows:
if not isinstance(row, dict):
raise HTTPException(502, "平台字段记录格式无法识别")
dataset = row.get("dataset")
owner = dataset.get("id") if isinstance(dataset, dict) else dataset
try:
owner = identifier(owner)
data = normalize(row, owner)
except Exception as exc:
from ..worldquant import WqError
if isinstance(exc, WqError):
raise HTTPException(502, str(exc)) from None
raise
if filters.dataset_id and owner != filters.dataset_id:
raise HTTPException(502, "平台返回其他数据集字段")
for remote_key, key in (
("instrumentType", "instrument_type"),
("instrument_type", "instrument_type"),
("region", "region"),
("universe", "universe"),
("delay", "delay"),
):
if remote_key in row and row[remote_key] != getattr(filters, key):
raise HTTPException(502, "平台字段范围与查询不一致")
items.append(
encode_snapshot(
dict(
**data,
field_id=data["id"],
dataset_id=owner,
dataset_name=label(dataset) or owner,
source="worldquant",
collection_version=None,
scope=filters.model_dump(include=set(Scope.model_fields)),
fetched_at=now(),
synced_at=None,
)
)
)
count = raw.get("count")
known_total = type(count) is int and count >= 0
more = (
bool(raw["next"])
if "next" in raw
else (filters.offset + len(items) < count if known_total else len(items) == filters.limit)
)
result = page(
items,
count if known_total else filters.offset + len(items) + int(more),
filters.limit,
filters.offset,
)
result.update(has_more=more, total_known=known_total)
return result
async def resolve_fields(self, scope, refs):
"""Resolve trusted source records before any collection member is written."""
if any(ref.scope.key() != scope.key() for ref in refs):
raise HTTPException(422, "不能跨区域、Top、Delay 或品种添加字段")
result = {}
for ref in refs:
if ref.source == "local":
dataset = await Catalog(self.db).dataset(scope, ref.dataset_id, lock=True)
if not dataset.field_version or dataset.field_version != ref.collection_version:
raise HTTPException(409, "字段来源已更新,请重新查询后添加")
entry = await self.db.get(CatalogEntry, (dataset.field_version, ref.field_id))
if not entry:
raise HTTPException(422, "字段不属于指定数据集")
scope_row = await self.db.get(CatalogScope, scope.key())
owner = await self.db.get(CatalogEntry, (scope_row.catalog_version, ref.dataset_id))
field = self.local_field(entry, ref.dataset_id, owner.name if owner else None, scope)
else:
offset, seen, field = 0, set(), None
while True:
response = await self.online_fields(
FieldFilters(
**scope.model_dump(),
q=ref.field_id,
dataset_id=ref.dataset_id,
limit=100,
offset=offset,
)
)
field = next((item for item in response["items"] if item["id"] == ref.field_id), None)
if field or not response["has_more"]:
break
ids = {item["id"] for item in response["items"]}
if not ids - seen:
raise HTTPException(502, "平台字段分页未前进")
seen.update(ids)
offset += len(response["items"])
if not field:
raise HTTPException(422, "在线字段已不可用,请重新查询")
if field["id"] in result and result[field["id"]]["dataset_id"] != field["dataset_id"]:
raise HTTPException(422, "同名字段的数据集归属冲突")
result[field["id"]] = field
return list(result.values())
async def get(self, preparation_id, version=None, lock=False):
query = select(DataPreparation).where(DataPreparation.id == preparation_id)
row = await self.db.scalar(query.with_for_update() if lock else query)
if not row:
raise HTTPException(404, "数据准备集合不存在")
if version is not None and row.version != version:
raise HTTPException(409, "集合已修改,请重新读取或选择;当前草稿已保留")
return row
async def output(self, row):
count, datasets = (
await self.db.execute(
select(func.count(), func.count(func.distinct(PreparationField.dataset_id))).where(
PreparationField.preparation_id == row.id
)
)
).one()
return encode_snapshot(
dict(
id=row.id,
name=row.name,
note=row.note,
scope=row.scope,
version=row.version,
field_count=count,
dataset_count=datasets,
created_at=row.created_at,
updated_at=row.updated_at,
)
)
async def list(self, q="", scope_key=None, limit=25, offset=0):
query = select(DataPreparation)
if scope_key:
query = query.where(DataPreparation.scope_key == scope_key)
if q:
query = query.where(
or_(
DataPreparation.name.ilike(contains(q), escape="\\"),
DataPreparation.note.ilike(contains(q), escape="\\"),
)
)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(DataPreparation.updated_at.desc(), DataPreparation.id).limit(limit).offset(offset)
)
return page([await self.output(row) for row in rows], total, limit, offset)
async def members(self, preparation_id, q="", dataset_id=None, limit=25, offset=0):
await self.get(preparation_id)
query = select(PreparationField).where(PreparationField.preparation_id == preparation_id)
if dataset_id:
query = query.where(PreparationField.dataset_id == dataset_id)
if q:
query = query.where(
or_(
PreparationField.field_id.ilike(contains(q), escape="\\"),
PreparationField.content["name"].as_string().ilike(contains(q), escape="\\"),
PreparationField.content["description"].as_string().ilike(contains(q), escape="\\"),
)
)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(PreparationField.dataset_id, PreparationField.field_id).limit(limit).offset(offset)
)
return page([row.content for row in rows], total, limit, offset)
async def create(self, name, note, scope, fields):
row = DataPreparation(
id=str(uuid4()), name=name, note=note, scope=scope.model_dump(), scope_key=scope.key()
)
self.db.add(row)
await self.db.flush()
await self.add(row, fields)
return await self.output(row)
async def add(self, row, fields):
for field in fields:
existing = await self.db.get(PreparationField, (row.id, field["id"]))
if existing:
if existing.dataset_id != field["dataset_id"]:
raise HTTPException(422, "同名字段的数据集归属冲突")
continue
self.db.add(
PreparationField(
preparation_id=row.id, field_id=field["id"], dataset_id=field["dataset_id"], content=field
)
)
await self.db.flush()
async def copy_dataset(self, body):
dataset = await Catalog(self.db).dataset(body.scope, body.dataset_id, lock=True)
if not dataset.field_version or dataset.field_version != body.collection_version:
raise HTTPException(409, "数据集尚未完整同步或版本已更新")
batch = await self.db.get(CatalogBatch, dataset.field_version)
if not batch.complete:
raise HTTPException(409, "数据集尚未完整同步")
source = await Catalog(self.db).detail(body.scope, body.dataset_id)
rows = await self.db.scalars(
select(CatalogEntry)
.where(CatalogEntry.batch_id == dataset.field_version)
.order_by(CatalogEntry.id)
)
fields = [self.local_field(row, body.dataset_id, source["name"], body.scope) for row in rows]
return await self.create(
f"{source['name'] or body.dataset_id} · {now():%Y%m%d-%H%M%S-%f}", "", body.scope, fields
)
async def remove(self, refs):
rows = [await self.get(ref.id, ref.version, lock=True) for ref in sorted(refs, key=lambda r: r.id)]
for row in rows:
await self.db.execute(delete(PreparationField).where(PreparationField.preparation_id == row.id))
await self.db.delete(row)
return {"deleted": len(rows)}
async def freeze(self, refs):
"""Lock collection versions and capture source-independent research snapshots atomically."""
snapshots = []
for ref in sorted(refs, key=lambda r: r.id):
row = await self.get(ref.id, ref.version, lock=True)
existing = await self.db.scalar(
select(ResearchInputSnapshot).where(
ResearchInputSnapshot.preparation_id == row.id,
ResearchInputSnapshot.preparation_version == row.version,
)
)
if existing:
snapshots.append(await self.snapshot(existing.id))
continue
fields = [
r.content
for r in await self.db.scalars(
select(PreparationField)
.where(PreparationField.preparation_id == row.id)
.order_by(PreparationField.field_id)
)
]
if not fields:
raise HTTPException(422, "空集合不能用于研究")
fixed = ResearchInputSnapshot(
id=str(uuid4()),
preparation_id=row.id,
preparation_version=row.version,
content=dict(
name=row.name,
scope=row.scope,
fields=fields,
field_ids=[f["id"] for f in fields],
field_types={f["id"]: f["field_type"] for f in fields},
dataset_ids=sorted({f["dataset_id"] for f in fields}),
),
)
self.db.add(fixed)
await self.db.flush()
snapshots.append(await self.snapshot(fixed.id))
return snapshots
async def snapshot(self, snapshot_id):
row = await self.db.get(ResearchInputSnapshot, snapshot_id)
if not row:
raise HTTPException(404, "研究输入快照不存在")
return encode_snapshot(
dict(
**row.content,
id=row.id,
preparation_id=row.preparation_id,
preparation_version=row.preparation_version,
created_at=row.created_at,
)
)
async def bind(self, body):
refs = getattr(body, "preparation_refs", [])
if refs:
fixed = await self.freeze(refs)
body.input_ids = list(dict.fromkeys([*body.input_ids, *[r["id"] for r in fixed]]))
body.preparation_refs = []
if not body.input_ids:
raise HTTPException(422, "请选择非空的数据准备集合")
limit = next(
m.max_length for m in type(body).model_fields["input_ids"].metadata if hasattr(m, "max_length")
)
if len(body.input_ids) > limit:
raise HTTPException(422, f"最多可选择 {limit} 个研究输入")
return body
+16 -2
View File
@@ -4,6 +4,8 @@ from pydantic import Field
from ..ai.alpha_tools import AlphaArgs
from ..ai.capabilities import Capability
from ..preparations.service import Preparations
from ..schemas import Contract
from .contracts import ChatboxResearchInput, InputPageArgs, ResearchInputSelection, ResearchPreviewInput
@@ -16,14 +18,26 @@ async def prepare(ctx, args):
return await ctx.business.research_builder.prepare(ResearchPreviewInput(**args.model_dump()))
INSTRUCTIONS = "Chatbox 是研究来源,不是 Alpha 备注。新候选由服务端记录会话和生成轮次;引用已有草稿和重跑保留原来源。\n数据集研究先用 search_catalog 读取实际字段/类型/集合版本,或 get_research_input 读取页面提供的固定输入。\n只有明确选出的字段才能 prepare_research_input;不要把一页搜索结果当作整集。页面 unsaved_field_selection 为 true 且无输入引用时,请用户保存选择并点击“用此输入研究”,不能忽略排除项。\n有输入快照时使用 prepare_research_backtest,提供研究假设、具名占位符和实际字段类型,模拟参数须匹配输入范围。模板中的所有数据字段使用绑定占位符。\n字段类型未知时说明限制;VECTOR 处理方式在模板中明确写出,不把 VECTOR 当作 MATRIX。绑定校验不等于 FASTEXPR 语义或平台权限验证。\n无数据集输入的直接表达式研究可以 prepare_backtest;不得声称来自已核验数据集。已有候选草稿先 get_backtest_draft,再按 ID 和版本引用。"
INSTRUCTIONS = "研究先用 search_data_preparations 查询可编辑集合,读取集合 ID 与 version 后使用 prepare_research_input 固定输入。已有快照使用 get_research_input。所有研究来源保留独立快照,删除集合不影响已有研究。构建回测需明确假设、字段绑定和范围;VECTOR 必须显式处理,不能当作 MATRIX。直接表达式回测不声明数据准备来源。"
class PreparationSearch(Contract):
q: str = Field(default="", max_length=300)
scope_key: str | None = None
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
CAPABILITIES = (
Capability(name="search_data_preparations", schema=PreparationSearch,
description="分页搜索数据准备集合,返回 ID、version、范围与字段数;非空集合可固定为研究输入。",
label="查询数据准备", renderer="catalog", effect="query",
handler=lambda ctx, args: Preparations(ctx.business.db).list(**args.model_dump())),
Capability(
name="prepare_research_input",
schema=ResearchInputSelection,
description="把明确选出的 1–100 个字段固定为研究输入快照;必须提供已读取的集合版本。只保存本地输入,不启动同步或回测。已有输入引用时直接读取,不重建。",
description="将数据准备集合的明确 ID 和 version 固定为研究快照,保留完整字段与数据集归属。已有快照直接读取。",
label="固定研究输入",
renderer="catalog",
effect="prepare",
+5 -1
View File
@@ -57,7 +57,11 @@ class Assets:
"view": ViewSpec,
"workflow": WorkflowSpec,
}[body.kind]
content = schema.model_validate(body.content).model_dump(mode="json")
parsed = schema.model_validate(body.content)
if body.kind == "feature":
from ..preparations.service import Preparations
await Preparations(self.db).bind(parsed)
content = parsed.model_dump(mode="json")
if body.kind == "workflow":
from .workflows import validate_graph
+3 -6
View File
@@ -5,16 +5,13 @@ from typing import Literal
from pydantic import Field, model_validator
from ..backtests.contracts import SimulationSettings, Source
from ..catalog.contracts import Scope
from ..preparations.contracts import PreparationReference
from ..schemas import Contract
from .expressions import PLACEHOLDER
class ResearchInputSelection(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
collection_version: str = Field(min_length=1, max_length=36)
field_ids: list[str] = Field(min_length=1, max_length=100)
items: list[PreparationReference] = Field(min_length=1, max_length=1)
class InputPageArgs(Contract):
@@ -50,7 +47,7 @@ class ChatboxResearchInput(Contract):
name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=2000)
template_input_id: str = Field(min_length=1, max_length=36)
input_snapshot_id: str = Field(min_length=1, max_length=36)
candidates: list[ResearchCandidate] = Field(min_length=1, max_length=100)
@model_validator(mode="after")
+9 -10
View File
@@ -11,7 +11,8 @@ from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Simulatio
from ..backtests.service import Backtests, uid
from ..catalog.research_metadata import ResearchMetadata
from ..catalog.service import Catalog
from ..models import Alpha, BacktestRun, CatalogResource, ResearchExperiment, TemplateInput
from ..models import Alpha, BacktestRun, CatalogResource, ResearchExperiment
from ..preparations.service import Preparations
from .assets import Assets
from .expressions import GROUPS, analyze, expand
from .serialization import encode_snapshot as jsonable_encoder
@@ -86,7 +87,7 @@ class Experiments:
"candidates": experiment["candidates"],
"hypothesis": experiment["hypothesis"],
"input_references": [
{k: entry[k] for k in ("id", "collection_version", "scope", "dataset_id")}
{k: entry[k] for k in ("id", "preparation_id", "preparation_version", "scope", "dataset_ids")}
for entry in experiment["inputs"]
],
"template_reference": {
@@ -133,6 +134,7 @@ class Experiments:
return validation
async def create(self, body, kind="template", extra_evidence=None, *, parent_snapshots=None):
await Preparations(self.db).bind(body)
asset = await self.assets.get(body.asset_id, body.version, "template") if body.asset_id else None
template = TemplateSpec.model_validate(asset["content"]) if asset else body.template
scope = scope_of(body.settings)
@@ -318,7 +320,8 @@ class Experiments:
kind=source_kind or experiment["kind"],
reference=reference or experiment_id,
research_id=experiment_id,
template_input_id=inputs[0]["id"] if len(inputs) == 1 else None,
input_snapshot_ids=[i["id"] for i in inputs],
input_snapshot_id=inputs[0]["id"] if len(inputs) == 1 else None,
hypothesis=experiment["hypothesis"][:2000],
),
candidates=[
@@ -342,6 +345,7 @@ class Experiments:
original = parents[0]
base = seed_settings(original["settings"])
expression = original["expression"]
await Preparations(self.db).bind(body)
snapshots, _ = await self.inputs(body.input_ids)
groups = defaultdict(list)
for snapshot in snapshots:
@@ -405,6 +409,7 @@ class Experiments:
)
async def generation_context(self, body):
await Preparations(self.db).bind(body)
snapshots, fields = await self.inputs(body.input_ids)
parents = await self.parents(body.parent_alpha_ids, body.parent_experiment_ids)
metadata = await ResearchMetadata(self.db).operators(limit=100)
@@ -414,7 +419,7 @@ class Experiments:
"hypothesis": body.hypothesis,
"method": body.method,
"inputs": [
{"id": item["id"], "scope": item["scope"], "dataset_id": item["dataset_id"]}
{"id": item["id"], "scope": item["scope"], "dataset_ids": item["dataset_ids"], "name": item["name"], "fields": item["fields"][:100]}
for item in snapshots
],
"fields": dict(list(fields.items())[:300]),
@@ -424,9 +429,3 @@ class Experiments:
],
"parents": [{**p, "candidates": p.get("candidates", [])[:10]} for p in parents],
}
async def available_inputs(self, limit=100):
rows = await self.db.scalars(
select(TemplateInput).order_by(TemplateInput.created_at.desc()).limit(limit)
)
return {"items": [await self.catalog.input(row.id) for row in rows]}
+2 -8
View File
@@ -28,12 +28,6 @@ from .workspace_contracts import (
router = APIRouter(prefix="/api/v1/research", tags=["research"], dependencies=[Depends(require_auth)])
@router.get("/inputs")
async def inputs(request: Request, limit: int = Query(100, ge=1, le=100)):
async with request.app.state.sessions() as db:
return await Experiments(db).available_inputs(limit)
@router.get("/assets")
async def assets(
request: Request,
@@ -96,10 +90,10 @@ async def import_commit(body: ImportCommit, request: Request):
@router.post("/generate", status_code=201)
async def generate(body: Generation, request: Request):
async with request.app.state.sessions() as db:
async with request.app.state.sessions.begin() as db:
context = await Experiments(db).generation_context(body)
result, evidence = await request_model(request.app.state.ai, context, OUTPUTS[body.method])
if body.method == "feature" and set(result.input_ids) != set(body.input_ids):
if body.method == "feature" and (result.preparation_refs or set(result.input_ids) != set(body.input_ids)):
raise HTTPException(422, "模型不能改变已固定的输入范围")
async with request.app.state.sessions.begin() as db:
asset = await Assets(db).save(
+2 -2
View File
@@ -476,9 +476,9 @@ class ResearchRuntime:
step = await db.get(ResearchStepRun, step_id)
if not step or step.status != "running":
return
if isinstance(result, FeatureSpec) and set(result.input_ids) != set(
if isinstance(result, FeatureSpec) and (result.preparation_refs or set(result.input_ids) != set(
[i["id"] for i in step.output["context"]["inputs"]]
):
)):
raise HTTPException(422, "模型不能改变已固定的输入范围")
# A paused/stopped run may collect this already-issued model output, but cannot advance.
asset = await Assets(db).save(
+10 -50
View File
@@ -5,12 +5,9 @@ not FASTEXPR operator semantics or the account's current platform permissions.
"""
from fastapi import HTTPException
from sqlalchemy import select
from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Source
from ..catalog.contracts import EntryOutput, InputPreparation
from ..catalog.service import Catalog
from ..models import CatalogEntry
from .expressions import analyze, expand
@@ -21,55 +18,18 @@ class ResearchBuilder:
self.backtests = backtests
async def select_input(self, body):
"""Fix explicit fields in one published version; reject missing or stale members."""
collection = await self.catalog.collection(body.scope, body.dataset_id)
chosen = set(body.field_ids)
if len(chosen) != len(body.field_ids) or not chosen.issubset(collection["field_ids"]):
raise HTTPException(422, "字段选择含重复、未知或其他数据集字段")
saved = await self.catalog.prepare(
InputPreparation(
scope=body.scope,
dataset_id=body.dataset_id,
collection_version=body.collection_version,
selection="explicit",
excluded_ids=[field for field in collection["field_ids"] if field not in chosen],
)
)
from ..preparations.service import Preparations
saved = (await Preparations(self.db).freeze(body.items))[0]
return await self.input_page(saved["id"])
async def input_page(self, input_id, limit=25, offset=0, q="", field_type=None):
"""Read the saved version, including field descriptions, with explicit pagination."""
saved = await self.catalog.input(input_id)
ids = [
field
for field in saved["field_ids"]
if q.lower() in field.lower()
and (field_type is None or saved["field_types"].get(field) == field_type)
]
page = ids[offset : offset + limit]
entries = {
row.id: row
for row in await self.db.scalars(
select(CatalogEntry).where(
CatalogEntry.batch_id == saved["collection_version"], CatalogEntry.id.in_(page)
)
)
}
return {
**{
k: saved[k]
for k in ("id", "scope", "dataset_id", "collection_version", "selection", "created_at")
},
"field_count": len(saved["field_ids"]),
"items": [
EntryOutput.model_validate(entries[field], from_attributes=True).model_dump()
for field in page
],
"total": len(ids),
"limit": limit,
"offset": offset,
"has_more": offset + limit < len(ids),
}
fields = [f for f in saved["fields"] if (not q or q.lower() in
" ".join(str(f.get(k) or "") for k in ("id", "name", "description", "dataset_id")).lower())
and (not field_type or f["field_type"] == field_type)]
return {**{k: v for k, v in saved.items() if k not in ("fields", "field_ids", "field_types")},
"field_count": len(saved["fields"]), "items": fields[offset:offset + limit],
"total": len(fields), "limit": limit, "offset": offset, "has_more": offset + limit < len(fields)}
async def prepare(self, body):
"""Bind templates against an immutable input, then reuse the fixed-preview interface.
@@ -77,7 +37,7 @@ class ResearchBuilder:
Raises HTTPException(422) for wrong scope, membership or declared type.
No expression execution or implicit cleaning/aggregation takes place here.
"""
saved = await self.catalog.input(body.template_input_id)
saved = await self.catalog.input(body.input_snapshot_id)
scope = saved["scope"]
candidates = []
for item in body.candidates:
@@ -114,7 +74,7 @@ class ResearchBuilder:
source = Source.model_validate(
{
**body.source.model_dump(),
"template_input_id": saved["id"],
"input_snapshot_id": saved["id"],
"hypothesis": body.hypothesis,
}
)
+2
View File
@@ -181,6 +181,8 @@ class Workflows:
for node in graph.nodes:
if node.type == "iterate" and node.config["max_rounds"] > body.budget.max_rounds:
raise HTTPException(422, "流程迭代上限超过本次授权轮数")
from ..preparations.service import Preparations
await Preparations(self.db).bind(body)
experiments = Experiments(self.db)
settings_variant = any(
n.type == "variant" and n.config.get("method") == "settings" for n in graph.nodes
+11 -5
View File
@@ -7,6 +7,7 @@ from pydantic import Field, field_validator, model_validator
from ..backtests.contracts import SimulationSettings
from ..catalog.contracts import Scope
from ..preparations.contracts import PreparationReference
from ..schemas import Contract
from .expressions import IDENTIFIER, PLACEHOLDER, normalize_template
@@ -70,7 +71,8 @@ class FeatureStep(Contract):
class FeatureSpec(Contract):
name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=10000)
input_ids: list[str] = Field(min_length=1, max_length=20)
input_ids: list[str] = Field(default_factory=list, max_length=20)
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
steps: list[FeatureStep] = Field(default_factory=list, max_length=30)
template: TemplateSpec | None = None
@@ -102,7 +104,8 @@ class Expansion(Contract):
asset_id: str | None = Field(default=None, max_length=36)
version: int | None = Field(default=None, ge=1)
template: TemplateSpec | None = None
input_ids: list[str] = Field(min_length=1, max_length=20)
input_ids: list[str] = Field(default_factory=list, max_length=20)
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
hypothesis: str = Field(min_length=1, max_length=10000)
settings: SimulationSettings
mode: Literal["all", "random"] = "all"
@@ -123,7 +126,8 @@ class Expansion(Contract):
class Generation(Contract):
name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=10000)
input_ids: list[str] = Field(min_length=1, max_length=20)
input_ids: list[str] = Field(default_factory=list, max_length=20)
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
parent_alpha_ids: list[str] = Field(default_factory=list, max_length=20)
parent_experiment_ids: list[str] = Field(default_factory=list, max_length=20)
method: Literal["template", "structure", "feature"] = "template"
@@ -131,7 +135,8 @@ class Generation(Contract):
class SettingVariants(Contract):
alpha_id: str = Field(min_length=1, max_length=100)
input_ids: list[str] = Field(min_length=1, max_length=100)
input_ids: list[str] = Field(default_factory=list, max_length=100)
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=100)
hypothesis: str = Field(default="保持表达式,比较字段共同支持的市场与设置", max_length=10000)
@@ -229,7 +234,8 @@ class FlowStart(Contract):
name: str = Field(min_length=1, max_length=200)
workflow_id: str | None = None
workflow_version: int | None = Field(default=None, ge=1)
input_ids: list[str] = Field(min_length=1, max_length=20)
input_ids: list[str] = Field(default_factory=list, max_length=20)
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
hypothesis: str = Field(min_length=1, max_length=10000)
settings: SimulationSettings
budget: Budget
+17
View File
@@ -7,6 +7,7 @@ from pydantic import Field, model_validator
from ..backtests.contracts import Candidate, SimulationSettings
from ..catalog.contracts import CatalogFilters, Scope
from ..preparations.contracts import PreparationReference
from ..research.workspace_contracts import TemplateSpec
from ..schemas import Contract
@@ -54,8 +55,24 @@ class Provenance(Contract):
parent_run_id: RunId | None = None
class PreparationSearch(Contract):
q: str = Field(default="", max_length=300)
scope_key: str | None = None
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class PreparationRead(Contract):
id: str = Field(min_length=1, max_length=36)
version: int = Field(ge=1)
q: str = Field(default="", max_length=300)
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class Submit(Contract):
name: str = Field(min_length=1, max_length=200)
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
candidates: list[DirectCandidate] = Field(min_length=1, max_length=100)
idempotency_key: Identifier
duplicate_policy: Literal["reject", "rerun"] = "reject"
+13 -1
View File
@@ -118,6 +118,17 @@ class ResearchAccess:
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"}),
@@ -298,7 +309,8 @@ class ResearchAccess:
# 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)
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"],