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"],
@@ -0,0 +1,75 @@
"""Replace saved dataset inputs with editable preparations and independent snapshots."""
import sqlalchemy as sa
from alembic import op
revision = "0015"
down_revision = "0014"
branch_labels = None
depends_on = None
def upgrade():
op.drop_table("template_inputs")
op.add_column("catalog_batches", sa.Column("job_id", sa.String(36), nullable=True))
op.add_column("catalog_batches", sa.Column("offset", sa.Integer(), nullable=False, server_default="0"))
op.execute("UPDATE catalog_batches SET job_id = id")
# Preserve unfinished catalog pagination; only legacy research inputs are discarded.
jobs = sa.table("sync_jobs", sa.column("id"), sa.column("checkpoint", sa.JSON()))
batches = sa.table("catalog_batches", sa.column("id"), sa.column("offset", sa.Integer()))
for job_id, checkpoint in op.get_bind().execute(sa.select(jobs.c.id, jobs.c.checkpoint)):
offset = (checkpoint or {}).get("offset", 0)
if type(offset) is int and offset >= 0:
op.get_bind().execute(batches.update().where(batches.c.id == job_id).values(offset=offset))
names = {"fk": "fk_%(table_name)s_%(column_0_name)s_%(referred_table_name)s"}
fk = next(
f
for f in sa.inspect(op.get_bind()).get_foreign_keys("catalog_batches")
if f["constrained_columns"] == ["id"]
)
with op.batch_alter_table("catalog_batches", naming_convention=names) as batch:
batch.drop_constraint(fk["name"] or "fk_catalog_batches_id_sync_jobs", type_="foreignkey")
batch.alter_column("job_id", existing_type=sa.String(36), nullable=False)
batch.create_foreign_key("fk_catalog_batches_job", "sync_jobs", ["job_id"], ["id"])
batch.create_index("ix_catalog_batches_job_id", ["job_id"])
op.create_table(
"data_preparations",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("name", sa.String(200), nullable=False),
sa.Column("note", sa.Text(), nullable=False),
sa.Column("scope_key", sa.String(200), nullable=False),
sa.Column("scope", sa.JSON(), nullable=False),
sa.Column("version", sa.Integer(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
)
op.create_index("ix_data_preparations_scope_key", "data_preparations", ["scope_key"])
op.create_table(
"preparation_fields",
sa.Column(
"preparation_id",
sa.String(36),
sa.ForeignKey("data_preparations.id", ondelete="CASCADE"),
primary_key=True,
),
sa.Column("field_id", sa.String(200), primary_key=True),
sa.Column("dataset_id", sa.String(200), nullable=False),
sa.Column("content", sa.JSON(), nullable=False),
)
op.create_index("ix_preparation_fields_dataset_id", "preparation_fields", ["dataset_id"])
op.create_table(
"research_input_snapshots",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("preparation_id", sa.String(36), nullable=False),
sa.Column("preparation_version", sa.Integer(), nullable=False),
sa.Column("content", sa.JSON(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.UniqueConstraint("preparation_id", "preparation_version"),
)
op.create_index(
"ix_research_input_snapshots_preparation_id", "research_input_snapshots", ["preparation_id"]
)
def downgrade():
raise RuntimeError("旧输入模型已移除;回退请恢复升级前数据库备份")
+3 -5
View File
@@ -43,8 +43,7 @@ if __name__ == "__main__":
"INSERT INTO research (alpha_id, note, tags, favorite, state, updated_at, version) VALUES ('MIGRATION_TEST', 'preserve research', '[]', false, 'inbox', now(), 7);"
)
)
command.upgrade(config, "head")
command.check(config)
command.upgrade(config, "0014")
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
("preserve research", 7)
]
@@ -56,7 +55,7 @@ if __name__ == "__main__":
("preserve research", 7)
]
print(
"PostgreSQL 17: 0002 → 0003, downgrade/re-upgrade, metadata check, Alpha/research preservation passed"
"PostgreSQL 17: catalog migration, 0014 downgrade/re-upgrade then head, metadata and Alpha preservation passed"
)
async def flow():
@@ -73,7 +72,6 @@ if __name__ == "__main__":
return httpx.Response(201, json={"token": {"expiry": 14400}})
if request.url.path == "/users/self":
return httpx.Response(200, json={"id": "PG_TEST_USER"})
assert request.method == "GET"
return catalog_response(request) or httpx.Response(404)
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
@@ -115,7 +113,7 @@ if __name__ == "__main__":
assert sorted(r.status_code for r in responses) == [200, 409]
await sync(catalog, "TEST_FIN")
assert (await prepare(client, version)).status_code == 409
persisted = (await client.get("/api/v1/catalog/inputs/" + draft["id"])).json()
persisted = (await client.get("/api/v1/research/input-snapshots/" + draft["id"])).json()
assert persisted == draft
print(
"PostgreSQL: real API/runner multi-page dedupe, immutable draft, refresh conflict and concurrent note CAS passed"
+139
View File
@@ -0,0 +1,139 @@
"""Isolated PostgreSQL migration and concurrency acceptance, using synthetic upstream only."""
import asyncio
import os
from alembic import command
from alembic.config import Config
from cryptography.fernet import Fernet
from sqlalchemy import text
from sqlalchemy.ext.asyncio import create_async_engine
URL = "postgresql+asyncpg://postgres:preparations-test-only@127.0.0.1:18437/preparations_test"
os.environ.update(
DATABASE_URL=URL,
ADMIN_PASSWORD="migration-test-only",
ENCRYPTION_KEY=Fernet.generate_key().decode(),
WQ_EMAIL="",
WQ_PASSWORD="",
)
async def sql(query):
engine = create_async_engine(URL)
try:
async with engine.begin() as db:
result = await db.execute(text(query))
return result.fetchall() if result.returns_rows else None
finally:
await engine.dispose()
async def acceptance():
import httpx
from app.catalog.contracts import CatalogJobInput, Scope
from app.catalog.service import Catalog
from app.config import Settings
from app.main import create_app
from app.worldquant import WqClient
from tests.catalog_fake import catalog_response
from tests.test_catalog import SCOPE, prepare, search, sync
def upstream(request):
if request.url.path == "/authentication":
return httpx.Response(201, json={"token": {"expiry": 14400}})
if request.url.path == "/users/self":
return httpx.Response(200, json={"id": "PG_TEST_USER"})
from tests.research_metadata_fake import response
metadata = response(request)
if metadata is not None:
return metadata
assert request.method == "GET"
return catalog_response(request) or httpx.Response(404)
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
app = create_app(settings, WqClient(settings, transport=httpx.MockTransport(upstream)))
async with app.router.lifespan_context(app):
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
headers={"X-WQ-Request": "1"},
) as client:
assert (
await client.post(
"/api/v1/auth/login", json={"username": "admin", "password": "migration-test-only"}
)
).status_code == 200
await client.put(
"/api/v1/account/credentials",
json={"email": "synthetic@example.com", "password": "synthetic-only"},
)
job = (await client.post("/api/v1/account/connect")).json()
await app.state.runner.execute(job["id"])
fixture = (client, app.state.runner, {})
await sync(fixture)
await sync(fixture, "TEST_FIN")
snapshot = (
await prepare(
client, (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"]
)
).json()
assert len(snapshot["field_ids"]) == 123
fields = await client.get("/api/v1/catalog/fields", params={**SCOPE, "category": "基本面", "limit": 2, "offset": 2})
assert fields.status_code == 200, fields.text
assert fields.json()["total"] == 123 and len(fields.json()["items"]) == 2
assert fields.json()["items"][0]["category"] == "基本面"
async def enqueue():
async with app.state.sessions.begin() as db:
return (await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)).id
jobs = await asyncio.gather(*(enqueue() for _ in range(5)))
assert len(set(jobs)) == 1, jobs
collection = (await client.get("/api/v1/data-preparations")).json()["items"][0]
ref = {"id": collection["id"], "version": collection["version"]}
async def freeze():
response = await client.post("/api/v1/data-preparations/freeze", json={"items": [ref]})
assert response.status_code == 201, response.text
return response.json()["items"][0]["id"]
assert len(set(await asyncio.gather(*(freeze() for _ in range(5))))) == 1
await client.delete(f"/api/v1/data-preparations/{ref['id']}?version={ref['version']}")
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
if __name__ == "__main__":
assert not asyncio.run(sql("SELECT tablename FROM pg_tables WHERE schemaname='public'")), (
"Requires an empty isolated database"
)
config = Config("alembic.ini")
command.upgrade(config, "0014")
# Existing catalog checkpoints survive the structural change; no research data migration.
asyncio.run(
sql(
"INSERT INTO sync_jobs(id,kind,status,payload,checkpoint,total,processed,failed,cancel_requested,created_at,updated_at) VALUES ('checkpoint-test','catalog_sync','failed','{}','{\"offset\": 100}',0,100,0,false,now(),now())"
)
)
asyncio.run(sql("INSERT INTO catalog_scopes(key,scope) VALUES ('checkpoint-scope','{}')"))
asyncio.run(
sql(
"INSERT INTO catalog_batches(id,scope_key,dataset_id,complete,count) VALUES ('checkpoint-test','checkpoint-scope',NULL,false,100)"
)
)
command.upgrade(config, "head")
command.check(config)
assert asyncio.run(
sql("SELECT job_id, catalog_batches.\"offset\" FROM catalog_batches WHERE id='checkpoint-test'")
) == [("checkpoint-test", 100)]
assert asyncio.run(sql("SELECT to_regclass('template_inputs')")) == [(None,)]
asyncio.run(sql("DELETE FROM catalog_batches WHERE id='checkpoint-test'"))
asyncio.run(sql("DELETE FROM catalog_scopes WHERE key='checkpoint-scope'"))
asyncio.run(sql("DELETE FROM sync_jobs WHERE id='checkpoint-test'"))
asyncio.run(acceptance())
print(
"PostgreSQL 17: 0014 → 0015 metadata, retained catalog checkpoint, concurrent job deduplication/freeze and independent snapshot passed"
)
+8 -24
View File
@@ -22,7 +22,7 @@ def research_step(text, returns, history):
return "get_backtest_results", {"run_id": run_id}
data = content(returns[-1])
return f"已读取本次研究的 {data['total']} 条真实保存结果;缺失 Sharpe 仍为未知。"
if context.get("unsaved_field_selection") and not context.get("template_input_id"):
if context.get("unsaved_field_selection") and not context.get("input_snapshot_id"):
return "请先保存字段选择,再点击用此输入研究。"
if not returns:
return "get_backtest_capabilities", {}
@@ -30,39 +30,23 @@ def research_step(text, returns, history):
data = content(last)
if "error" in data:
return f"研究尚未完成:{data['error']}"
scope = context.get("catalog_scope") or {
"instrument_type": "EQUITY",
"region": "USA",
"universe": "TOP3000",
"delay": 1,
}
if last.tool_name == "get_backtest_capabilities":
if context.get("template_input_id"):
if context.get("input_snapshot_id"):
return "get_research_input", {
"input_id": context["template_input_id"],
"input_id": context["input_snapshot_id"],
"field_type": "MATRIX",
"limit": 1,
}
return "search_catalog", {"filters": {**scope, "q": "TEST_FIN", "limit": 1}}
if last.tool_name == "search_catalog":
if data["dataset_id"] is None:
return "search_catalog", {
"dataset_id": data["items"][0]["id"],
"filters": {**scope, "field_type": "MATRIX", "limit": 1},
}
return "prepare_research_input", {
"scope": scope,
"dataset_id": data["dataset_id"],
"collection_version": data["collection_version"],
"field_ids": [data["items"][0]["id"]],
}
return "search_data_preparations", {"limit": 1}
if last.tool_name == "search_data_preparations":
return "prepare_research_input", {"items": [{"id": data["items"][0]["id"], "version": data["items"][0]["version"]}]}
if last.tool_name in ("get_research_input", "prepare_research_input"):
field = data["items"][0]
field = next(f for f in data["items"] if f["field_type"] == "MATRIX")
saved_scope = data["scope"]
return "prepare_research_backtest", {
"name": "Chatbox 数据集研究",
"hypothesis": "验证所选合成字段的横截面排序信号",
"template_input_id": data["id"],
"input_snapshot_id": data["id"],
"candidates": [
{
"client_item_id": "research-1",
+2 -2
View File
@@ -30,7 +30,7 @@ async def acceptance():
from app.config import Settings
from app.main import create_app
from app.models import TemplateInput
from app.models import ResearchInputSnapshot
from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure
from tests.test_backtests import setup
@@ -66,7 +66,7 @@ async def acceptance():
}
async with app.state.sessions() as db:
fixed = await db.scalar(select(TemplateInput))
fixed = await db.scalar(select(ResearchInputSnapshot))
body = {
"request_id": "finite-run",
"name": "PG 有限研究",
+3 -2
View File
@@ -170,14 +170,15 @@ async def main(args):
fields = await api("GET", f"/catalog/datasets/pv1/fields?{params}&limit=100")
fixed_input = await api(
"POST",
"/catalog/inputs",
"/data-preparations/from-dataset",
{
"scope": scope,
"dataset_id": "pv1",
"collection_version": fields["collection_version"],
"selection": "all",
},
)
fixed_input = (await api("POST", "/data-preparations/freeze", {
"items": [{"id": fixed_input["id"], "version": fixed_input["version"]}]}))["items"][0]
availability = await api(
"POST", "/catalog/field-availability/refresh", {"field_id": "close", "scope": scope}
)
+2 -2
View File
@@ -28,7 +28,7 @@ async def acceptance():
from app.config import Settings
from app.main import create_app
from app.models import ResearchExperiment, ResearchParent, TemplateInput
from app.models import ResearchExperiment, ResearchInputSnapshot, ResearchParent
from tests.test_research_outcomes import (
test_feature_conversion_keeps_original_version_through_experiment,
test_lineage_retains_multiple_parents_and_descendants,
@@ -46,7 +46,7 @@ async def acceptance():
)
assert response.status_code == 200
async with app.state.sessions() as db:
fixed = await db.scalar(select(TemplateInput))
fixed = await db.scalar(select(ResearchInputSnapshot))
for experiment in await db.scalars(select(ResearchExperiment)):
for parent in experiment.parents:
assert await db.get(ResearchParent, (experiment.id, parent["kind"], parent["id"]))
+2 -2
View File
@@ -30,7 +30,7 @@ async def acceptance():
from app.config import Settings
from app.main import create_app
from app.models import TemplateInput
from app.models import ResearchInputSnapshot
from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure
from tests.test_backtests import setup
@@ -67,7 +67,7 @@ async def acceptance():
}
async with app.state.sessions() as db:
fixed = await db.scalar(select(TemplateInput))
fixed = await db.scalar(select(ResearchInputSnapshot))
body = {
"request_id": "finite-run",
"name": "PG 有限研究",
+6 -9
View File
@@ -10,11 +10,10 @@ from sqlalchemy import func, select
from app.ai.capabilities import ToolContext, assemble
from app.ai.tools import CAPABILITIES
from app.business import Business
from app.models import AIMessage, AIRun, AIToolCall, Alpha, Research, TemplateInput
from app.models import AIMessage, AIRun, AIToolCall, Alpha, Research, ResearchInputSnapshot
from app.research.service import ResearchBuilder
from tests.test_ai import configure, single_tool_factory, start
from tests.test_api import seed
from tests.test_catalog import SCOPE
from tests.test_catalog import catalog as catalog_fixture
from tests.test_research_integration import fixed_input as fixed_input_fixture
@@ -62,15 +61,13 @@ async def test_prepare_failure_rolls_back_artifact_but_keeps_audit(app, logged_i
# select_input has already persisted the new input before requesting its result page.
raise HTTPException(422, "准备输入后的校验失败")
updated = await logged_in.patch(f"/api/v1/data-preparations/{fixed_input['preparation_id']}",
json={"version": fixed_input["preparation_version"], "name": "new version"})
assert updated.status_code == 200
monkeypatch.setattr(ResearchBuilder, "input_page", unavailable_page)
app.state.ai.model_factory = single_tool_factory(
"prepare_research_input",
{
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
{"items": [{"id": fixed_input["preparation_id"], "version": updated.json()["version"]}]},
)
_, run, _ = await start(app, logged_in, "保存研究输入")
call = run["tools"][0]
@@ -78,7 +75,7 @@ async def test_prepare_failure_rolls_back_artifact_but_keeps_audit(app, logged_i
assert call["presentation"]["effect"] == "prepare"
assert call["result"]["error"] == "准备输入后的校验失败"
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(TemplateInput)) == 1
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 1
assert (await db.get(AIToolCall, call["id"])).status == "failed"
+3 -2
View File
@@ -212,8 +212,7 @@ def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, m
},
)
for _ in range(2):
command.upgrade(config, "head")
command.check(config)
command.upgrade(config, "0014")
alphas = sa.Table("alphas", sa.MetaData(), autoload_with=engine)
with engine.connect() as db:
rows = db.execute(sa.select(alphas).order_by(alphas.c.id)).mappings().all()
@@ -224,4 +223,6 @@ def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, m
record = db.execute(sa.select(research)).mappings().one()
assert record["tags"] == ["PPAC"] and record["note"] == "keep" and record["version"] == 7
command.downgrade(config, "0010")
command.upgrade(config, "head")
command.check(config)
engine.dispose()
+1
View File
@@ -203,6 +203,7 @@ async def test_export_formula_injection_and_detail_variants(app, logged_in):
assert row["name"] == "'=DANGEROUS()" and row["note"] == "' @formula()"
assert (await logged_in.get(f"{PREFIX}/alphas/super1/pnl")).json() == {
"cached": False,
"series": [],
"points": [],
"fetched_at": None,
}
+21 -13
View File
@@ -86,16 +86,24 @@ async def search(client, suffix="/datasets", **params):
async def prepare(client, version, **changes):
return await client.post(
BASE + "/inputs",
json={
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": version,
"selection": "all",
**changes,
},
)
response = await client.post("/api/v1/data-preparations/from-dataset", json={
"scope": changes.get("scope", SCOPE), "dataset_id": changes.get("dataset_id", "TEST_FIN"),
"collection_version": version,
})
if response.status_code != 201:
return response
collection = response.json()
if changes.get("excluded_ids"):
response = await client.patch(f"/api/v1/data-preparations/{collection['id']}/fields", json={
"version": collection["version"], "remove_ids": changes["excluded_ids"]})
if response.status_code != 200:
return response
collection = response.json()
response = await client.post("/api/v1/data-preparations/freeze", json={
"items": [{"id": collection["id"], "version": collection["version"]}]})
if response.status_code != 201:
return response
return httpx.Response(201, json=response.json()["items"][0])
async def test_complete_workflow_filters_notes_immutable_input(catalog):
@@ -112,7 +120,7 @@ async def test_complete_workflow_filters_notes_immutable_input(catalog):
response = await prepare(client, version)
assert response.status_code == 201, response.text
draft = response.json()
assert len(draft["field_ids"]) == 123 and draft["status"] == "draft"
assert len(draft["field_ids"]) == 123 and draft["preparation_version"] == 1
for suffix in ["/datasets/TEST_FIN", "/datasets/TEST_FIN/fields/TEST_FIN_122"]:
detail = await search(client, suffix)
assert detail["research"]["version"] == 1
@@ -135,7 +143,7 @@ async def test_complete_workflow_filters_notes_immutable_input(catalog):
assert newer["collection_version"] != version
assert (await prepare(client, version)).status_code == 409
assert len((await prepare(client, newer["collection_version"])).json()["field_ids"]) == 125
assert (await client.get(BASE + "/inputs/" + draft["id"])).json() == draft
assert (await client.get("/api/v1/research/input-snapshots/" + draft["id"])).json() == draft
assert (await search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122"))["research"][
"note"
] == "保留研究备注"
@@ -219,7 +227,7 @@ async def test_scope_ownership_empty_and_unknown_fields_are_rejected(catalog):
await sync(catalog, scope=other)
await sync(catalog, "TEST_FIN", scope=other)
assert (await prepare(client, version, scope=other)).status_code == 409
assert len((await client.get(BASE + "/inputs", params=SCOPE)).json()) == 1
assert (await client.get(BASE + "/inputs", params=SCOPE)).status_code == 404
async def test_catalog_authentication_and_origin(app, client):
+26 -1
View File
@@ -162,7 +162,8 @@ async def test_official_sdk_client_and_error_contract(mcp_app):
async with ClientSession(streams[0], streams[1]) as client:
await client.initialize()
listed = await client.list_tools()
assert len(listed.tools) == 18
assert len(listed.tools) == 20
assert {"search_data_preparations", "get_data_preparation"} <= {t.name for t in listed.tools}
caps = await client.call_tool("get_research_capabilities", {})
assert caps.structured_content["max_candidates"] == 100
result = await client.call_tool("submit_backtests", submission())
@@ -459,3 +460,27 @@ async def test_worldquant_authentication_permissions_and_challenge(mcp_app):
assert missing.structured_content["error"]["code"] == "CREDENTIALS_NOT_CONFIGURED"
wrong = await mcp_app.state.mcp.invoke(principal, "get_worldquant_connection", {"job_id": "unrelated"})
assert wrong.structured_content["error"]["code"] == "NOT_FOUND"
async def test_mcp_preparations_freeze_at_submit_and_survive_deletion(mcp_app):
from app.catalog.contracts import Scope
from app.preparations.contracts import PreparationReference
from app.preparations.service import Preparations
principal, _ = await credentials(mcp_app)
async with mcp_app.state.sessions.begin() as db:
collection = await Preparations(db).create("MCP prepared", "", Scope(region="USA", universe="TOP3000", delay=1),
[{"id": "close", "field_id": "close", "name": "Close", "dataset_id": "pv1", "dataset_name": "Price",
"description": "Synthetic close", "field_type": "MATRIX", "source": "local", "fetched_at": now().isoformat()}])
refs = [{"id": collection["id"], "version": 1}]
found = await invoke(mcp_app, principal, "search_data_preparations", {"q": "MCP prepared"})
assert found["items"][0]["id"] == collection["id"]
detail = await invoke(mcp_app, principal, "get_data_preparation", {**refs[0], "limit": 1})
assert detail["fields"]["items"][0]["dataset_id"] == "pv1"
result = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
async with mcp_app.state.sessions.begin() as db:
run = await db.get(BacktestRun, result["backtest_run_id"])
snapshot_id = run.source["input_snapshot_ids"][0]
await Preparations(db).remove([PreparationReference(**refs[0])])
assert (await Preparations(db).snapshot(snapshot_id))["fields"][0]["description"] == "Synthetic close"
# Idempotent replay uses the already fixed run even after the collection is gone.
replay = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
assert replay["backtest_run_id"] == result["backtest_run_id"]
+352
View File
@@ -0,0 +1,352 @@
"""Preparation contracts across real HTTP, transactions and immutable research inputs."""
from argparse import Namespace
import pytest
from sqlalchemy import select
from app.catalog.contracts import CatalogJobInput, Scope
from app.catalog.service import Catalog
from app.models import CatalogBatch, Job, ResearchInputSnapshot
from tests.catalog_fake import field_records
from tests.test_catalog import SCOPE, sync
from tests.test_catalog import catalog as catalog_fixture
catalog = catalog_fixture
BASE = "/api/v1/data-preparations"
def reference(field):
return {k: field[k] for k in ("scope", "dataset_id", "field_id", "source", "collection_version")}
async def copied(catalog):
client = catalog[0]
await sync(catalog)
job = await sync(catalog, "TEST_FIN")
response = await client.post(
BASE + "/from-dataset",
json={"scope": SCOPE, "dataset_id": "TEST_FIN", "collection_version": job["id"]},
)
assert response.status_code == 201, response.text
return response.json()
async def test_multi_dataset_collection_atomic_changes_and_snapshot_independence(catalog):
client, runner, state = catalog
collection = await copied(catalog)
state["fields"] = field_records("TEST_NEWS", 3)
await sync(catalog, "TEST_NEWS")
query = (await client.get("/api/v1/catalog/fields", params={**SCOPE, "dataset_id": "TEST_NEWS"})).json()
assert query["total"] == 3
ref = reference(query["items"][0])
response = await client.patch(
f"{BASE}/{collection['id']}/fields", json={"version": 1, "fields": [ref, ref]}
)
assert response.status_code == 200, response.text
collection = response.json()
assert collection["field_count"] == 124 and collection["dataset_count"] == 2
bad = {**ref, "scope": {**SCOPE, "delay": 0}}
response = await client.patch(
f"{BASE}/{collection['id']}/fields",
json={"version": 2, "fields": [ref, bad], "remove_ids": ["TEST_FIN_001"]},
)
assert response.status_code == 422
assert (await client.get(f"{BASE}/{collection['id']}")).json()["version"] == 2
snapshot = (
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
).json()["items"][0]
assert snapshot["dataset_ids"] == ["TEST_FIN", "TEST_NEWS"]
assert all("description" in f and f["dataset_id"] for f in snapshot["fields"])
assert (
await client.patch(f"{BASE}/{collection['id']}", json={"version": 2, "name": "edited"})
).status_code == 200
assert (
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
).status_code == 409
assert (await client.delete(f"{BASE}/{collection['id']}?version=3")).status_code == 200
assert (await client.get(f"{BASE}/{collection['id']}")).status_code == 404
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
async with runner.sessions() as db:
assert await db.get(ResearchInputSnapshot, snapshot["id"])
@pytest.mark.parametrize(
"key,value", [("region", "CHN"), ("universe", "TOP1000"), ("delay", 0), ("instrument_type", "FUTURE")]
)
async def test_scope_dimensions_rejected_without_creating_collection(catalog, key, value):
client = catalog[0]
await copied(catalog)
field = (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["items"][0]
before = (await client.get(BASE)).json()["total"]
response = await client.post(
BASE, json={"name": "bad", "scope": {**SCOPE, key: value}, "fields": [reference(field)]}
)
assert response.status_code == 422
assert (await client.get(BASE)).json()["total"] == before
async def test_online_fields_without_sync_and_local_search_before_pagination(catalog):
client, _, _ = catalog
response = await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100})
assert response.status_code == 200, response.text
first = response.json()
assert len(first["items"]) == 100 and first["has_more"]
second = (
await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100, "offset": 100})
).json()
assert len(second["items"]) > 0
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
field = reference(second["items"][-1])
response = await client.post(BASE, json={"name": "online only", "scope": SCOPE, "fields": [field]})
assert response.status_code == 201, response.text
assert response.json()["field_count"] == 1
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
await copied(catalog)
query = (
await client.get(
"/api/v1/catalog/fields", params={**SCOPE, "q": "合成字段说明 11", "limit": 2, "offset": 2}
)
).json()
assert query["total"] == 11 and len(query["items"]) == 2
response = await client.get(
"/api/v1/catalog/fields", params={**SCOPE, "coverage_min": 0.9, "coverage_max": 0.5}
)
assert response.status_code == 422
async def test_empty_crud_copy_and_batch_delete_are_version_checked(catalog):
client = catalog[0]
first = (await client.post(BASE, json={"name": "empty", "scope": SCOPE})).json()
assert first["field_count"] == 0
assert (
await client.post(BASE + "/freeze", json={"items": [{"id": first["id"], "version": 1}]})
).status_code == 422
second = (await client.post(f"{BASE}/{first['id']}/copy", json={"version": 1})).json()
response = await client.post(
BASE + "/batch-delete",
json={"items": [{"id": first["id"], "version": 1}, {"id": second["id"], "version": 99}]},
)
assert response.status_code == 409
assert (await client.get(BASE)).json()["total"] == 2
assert (
await client.patch(
f"{BASE}/{first['id']}", json={"name": "x", "version": 1, "scope": {**SCOPE, "delay": 0}}
)
).status_code == 422
async def test_full_sync_partial_failure_and_resume_publish_per_dataset(catalog):
client, runner, state = catalog
await copied(catalog)
async with runner.sessions.begin() as db:
service = Catalog(db)
job = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
duplicate = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
assert duplicate.id == job.id
# The fixture deliberately returns TEST_FIN ownership for the other datasets.
await runner.execute(job.id)
failed = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert failed["status"] == "completed_with_errors" and failed["processed"] == 1
version = (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
"collection_version"
]
assert version != job.id
async with runner.sessions() as db:
assert (await db.get(CatalogBatch, version)).job_id == job.id
state["fields"] = field_records("TEST_NEWS", 3)
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
await runner.execute(job.id)
retried = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert retried["processed"] == 2 and retried["failed"] == 1
assert (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
"collection_version"
] == version
state["fields"] = field_records("TEST_UNKNOWN", 0)
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
await runner.execute(job.id)
done = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert done["status"] == "completed" and done["processed"] == 3 and done["failed"] == 0
async def test_cli_timeout_reuses_active_job_and_does_not_cancel(catalog, monkeypatch, capsys):
from app import cli
_, runner, _ = catalog
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
args = Namespace(
region="USA",
universe="TOP3000",
delay=1,
instrument_type="EQUITY",
resume_job=None,
wait_timeout=0.01,
)
assert await cli.catalog_sync_command(args) == 4
assert await cli.catalog_sync_command(args) == 4
async with runner.sessions() as db:
rows = list(await db.scalars(select(Job).where(Job.kind == "catalog_full_sync")))
assert len(rows) == 1 and rows[0].status == "queued" and not rows[0].cancel_requested
assert "等待超时" in capsys.readouterr().out
async def test_collection_version_checked_at_research_submission(catalog):
client = catalog[0]
collection = await copied(catalog)
response = await client.patch(f"{BASE}/{collection['id']}", json={"version": 1, "name": "changed"})
assert response.status_code == 200
body = {
"inline": {
"name": "prepared backtest",
"preparation_refs": [{"id": collection["id"], "version": 1}],
"candidates": [
{
"client_item_id": "one",
"expression": "rank(TEST_FIN_001)",
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
}
],
}
}
assert (await client.post("/api/v1/backtests/previews", json=body)).status_code == 409
body["inline"]["preparation_refs"][0]["version"] = 2
response = await client.post("/api/v1/backtests/previews", json=body)
assert response.status_code == 201, response.text
snapshot_id = response.json()["source"]["input_snapshot_ids"][0]
await client.delete(f"{BASE}/{collection['id']}?version=2")
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot_id}")).json()["fields"][1][
"dataset_id"
] == "TEST_FIN"
@pytest.mark.parametrize(
"status,code",
[
("completed", 0),
("completed_with_errors", 1),
("failed", 1),
("cancelled", 1),
("waiting_connection", 3),
("waiting_auth", 3),
],
)
async def test_cli_terminal_exit_codes(catalog, monkeypatch, status, code):
from app import cli
_, runner, _ = catalog
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
original = Catalog.create_job
async def terminal(self, body, **kwargs):
result = await original(self, body, **kwargs)
row = await self.db.get(Job, result.id)
row.status = status
return result
monkeypatch.setattr(Catalog, "create_job", terminal)
args = Namespace(
region="USA", universe="TOP3000", delay=1, instrument_type="EQUITY", resume_job=None, wait_timeout=1
)
assert await cli.catalog_sync_command(args) == code
async def test_full_sync_restart_keeps_page_and_auth_pauses_all_datasets(catalog, monkeypatch):
import asyncio
from app.worldquant import WqError
client, runner, _ = catalog
async with runner.sessions.begin() as db:
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
original = runner.client.catalog_page
interrupted = False
async def pages(scope, dataset_id, offset):
nonlocal interrupted
if dataset_id == "TEST_FIN" and offset == 100 and not interrupted:
interrupted = True
runner.stopping = True
raise asyncio.CancelledError()
if dataset_id == "TEST_NEWS":
raise WqError("synthetic network interruption", "network_error")
return await original(scope, dataset_id, offset)
monkeypatch.setattr(runner.client, "catalog_page", pages)
await runner.execute(job.id)
async with runner.sessions() as db:
batch = await db.scalar(
select(CatalogBatch).where(CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_FIN")
)
assert batch.offset == 100 and not batch.complete
assert (await db.get(Job, job.id)).status == "queued"
runner.stopping = False
await runner.execute(job.id)
row = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert row["status"] == "waiting_connection" and row["processed"] == 1
assert row["checkpoint"]["dataset_id"] == "TEST_NEWS"
assert row["checkpoint"]["datasets_completed"] == 1
async with runner.sessions() as db:
assert not await db.scalar(
select(CatalogBatch).where(
CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_UNKNOWN"
)
)
await client.post(f"/api/v1/sync-jobs/{job.id}/cancel")
assert (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()["status"] == "cancelled"
async def test_online_instrument_type_is_verified(catalog):
client, _, state = catalog
state["fields"][0]["instrumentType"] = "FUTURE"
assert (await client.get("/api/v1/catalog/worldquant/fields", params=SCOPE)).status_code == 502
async def test_mcp_collection_reads_and_versioned_submit_contract(catalog):
from types import SimpleNamespace
from app.research_access.contracts import PreparationRead, PreparationSearch, Submit
from app.research_access.service import ResearchAccess
client, runner, _ = catalog
collection = await copied(catalog)
async with runner.sessions.begin() as db:
access = ResearchAccess(
db, SimpleNamespace(token_id="test", admin_id=1), runner.client, "http://testserver"
)
result = await access.preparations(PreparationSearch(q=collection["name"]))
assert result["items"][0]["version"] == 1
detail = await access.preparation(PreparationRead(id=collection["id"], version=1, limit=1))
assert detail["fields"]["has_more"] and detail["fields"]["items"][0]["dataset_id"] == "TEST_FIN"
# Schema accepts versioned references and rejects accidental old dataset input payloads.
from tests.test_mcp import submission
payload = submission()
payload["preparation_refs"] = [{"id": collection["id"], "version": 1}]
assert Submit.model_validate(payload).preparation_refs[0].id == collection["id"]
async def test_retry_full_job_reuses_another_active_job(catalog):
from app.business import Business
_, runner, _ = catalog
body = CatalogJobInput(scope=Scope(**SCOPE))
async with runner.sessions.begin() as db:
old = await Catalog(db).create_job(body, full=True)
(await db.get(Job, old.id)).status = "failed"
async with runner.sessions.begin() as db:
new = await Catalog(db).create_job(body, full=True)
assert new.id != old.id
async with runner.sessions.begin() as db:
result = await Business(db).retry_job(old.id)
assert result["id"] == new.id
assert (await db.get(Job, old.id)).status == "failed"
async def test_retry_waiting_full_job_requeues_its_checkpoint(catalog):
from app.business import Business
_, runner, _ = catalog
async with runner.sessions.begin() as db:
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
row = await db.get(Job, job.id)
row.status, row.checkpoint = "waiting_connection", {"offset": 100}
async with runner.sessions.begin() as db:
result = await Business(db).retry_job(job.id)
assert result["status"] == "queued" and result["checkpoint"]["offset"] == 100
+12 -24
View File
@@ -11,7 +11,7 @@ from app.ai.tools import CAPABILITIES
from app.alphas import upsert_alpha
from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput
from app.business import Business
from app.models import BacktestPreview, BacktestRun, Research, TemplateInput
from app.models import BacktestPreview, BacktestRun, Research, ResearchInputSnapshot
from tests.test_ai import configure, single_tool_factory
from tests.test_backtests import execute, setup, start
from tests.test_catalog import SCOPE, prepare, sync
@@ -47,7 +47,7 @@ def construction(input_id):
return {
"name": "字段研究",
"hypothesis": "显式字段排序",
"template_input_id": input_id,
"input_snapshot_id": input_id,
"candidates": [
{
"client_item_id": "one",
@@ -66,7 +66,7 @@ async def test_chatbox_catalog_to_results_and_alpha_sources(app, logged_in, fixe
conversation = (await logged_in.post("/api/v1/ai/conversations")).json()["id"]
context = {"page": "datasets", "catalog_scope": SCOPE, "dataset_id": "TEST_FIN"}
if use_saved_input:
context["template_input_id"] = fixed_input["id"]
context["input_snapshot_id"] = fixed_input["id"]
run = await ask(logged_in, conversation, "研究此输入" if use_saved_input else "自行选字段研究", context)
assert run["status"] == "waiting_approval", run
assert not platform.posts
@@ -75,7 +75,7 @@ async def test_chatbox_catalog_to_results_and_alpha_sources(app, logged_in, fixe
assert source["kind"] == "chatbox"
assert source["reference"] == conversation
assert source["research_id"] == run["id"]
assert source["template_input_id"]
assert source["input_snapshot_id"]
assert approval["preview"]["backtest"]["items"][0]["expression"] == "rank(TEST_FIN_001)"
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
@@ -128,7 +128,7 @@ async def test_invalid_construction_has_no_partial_preview(logged_in, app, fixed
elif invalid == "unknown_type":
item["bindings"]["signal"] = {"field_id": "TEST_FIN_122", "field_type": "MATRIX"}
else:
body["template_input_id"] = "missing"
body["input_snapshot_id"] = "missing"
response = await logged_in.post("/api/v1/backtests/research-previews", json=body)
assert response.status_code == (404 if invalid == "missing_input" else 422), response.text
async with app.state.sessions() as db:
@@ -144,15 +144,10 @@ async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in
assert page["field_count"] == page["total"] == 123 and len(page["items"]) == 23
assert page["items"][-1]["field_type"] == "FUTURE_TYPE" and not page["has_more"]
assert page["_meta"]["source"] == "local_database"
selected = await tool(
"prepare_research_input",
{
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
)
subset = (await logged_in.post("/api/v1/data-preparations", json={"name": "one field", "scope": SCOPE,
"fields": [{"scope": SCOPE, "dataset_id": "TEST_FIN", "field_id": "TEST_FIN_001", "source": "local",
"collection_version": fixed_input["fields"][0]["collection_version"]}]})).json()
selected = await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
assert selected["field_count"] == 1
bad = construction(selected["id"])
bad["candidates"][0]["bindings"]["signal"]["field_id"] = "TEST_FIN_002"
@@ -166,19 +161,12 @@ async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in
"/api/v1/backtests/research-previews", json=construction(fixed_input["id"])
)
assert response.status_code == 201, response.text
await logged_in.patch(f"/api/v1/data-preparations/{subset['id']}", json={"name": "edited", "version": subset["version"]})
with pytest.raises(HTTPException) as exc:
await tool(
"prepare_research_input",
{
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
)
await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
assert exc.value.status_code == 409
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(TemplateInput)) == 2
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 2
async def test_multiple_origins_preserve_research_and_do_not_duplicate_alphas(app, logged_in):
+20 -8
View File
@@ -213,7 +213,7 @@ async def test_operator_annotation_and_refresh_preserve_local(app, logged_in, re
async def test_settings_variant_requires_all_fields_in_target(app, logged_in, research_input, catalog):
from app.models import CatalogScope, TemplateInput
from app.models import CatalogScope, ResearchInputSnapshot
from app.research.experiments import Experiments
from app.research.workspace_contracts import SettingVariants
@@ -232,14 +232,13 @@ async def test_settings_variant_requires_all_fields_in_target(app, logged_in, re
db.add(CatalogScope(key=target_key, scope=target_scope))
await db.flush()
db.add(
TemplateInput(
ResearchInputSnapshot(
id="target",
scope_key=target_key,
dataset_id="TEST_FIN",
collection_version=research_input["collection_version"],
selection="explicit",
field_ids=["TEST_FIN_001"],
field_types={"TEST_FIN_001": "MATRIX"},
preparation_id="target-preparation",
preparation_version=1,
content={**{k: v for k, v in research_input.items() if k not in ("id", "preparation_id", "preparation_version", "created_at")}, "scope": target_scope,
"field_ids": ["TEST_FIN_001"], "field_types": {"TEST_FIN_001": "MATRIX"},
"fields": [f for f in research_input["fields"] if f["id"] == "TEST_FIN_001"]},
)
)
await db.flush()
@@ -514,3 +513,16 @@ def test_real_seed_settings_preserve_execution_options_and_reject_unknowns():
assert snapshot["startDate"] == "2014-01-01"
with pytest.raises(ValidationError):
seed_settings({**snapshot, "unknownOption": True})
async def test_model_cannot_append_preparation_references_to_fixed_inputs(app, logged_in, research_input, monkeypatch):
from app.models import ResearchAsset
from app.research import routes
from app.research.workspace_contracts import FeatureSpec
async def model(*args):
return FeatureSpec(name="untrusted", hypothesis="test", input_ids=[research_input["id"]],
preparation_refs=[{"id": research_input["preparation_id"], "version": research_input["preparation_version"]}]), {}
monkeypatch.setattr(routes, "request_model", model)
response = await logged_in.post("/api/v1/research/generate", json={"name": "test", "hypothesis": "test", "method": "feature", "input_ids": [research_input["id"]]})
assert response.status_code == 422 and "不能改变" in response.text
async with app.state.sessions() as db:
assert not await db.scalar(select(ResearchAsset).where(ResearchAsset.name == "untrusted"))