refactor: unify data preparations and research input snapshots
Deploy production / deploy (push) Successful in 53s
Deploy production / deploy (push) Successful in 53s
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"] = {
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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,
|
||||
},
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Data preparation collections and immutable research inputs."""
|
||||
@@ -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
|
||||
@@ -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},
|
||||
}
|
||||
@@ -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
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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]}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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("旧输入模型已移除;回退请恢复升级前数据库备份")
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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",
|
||||
|
||||
@@ -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 有限研究",
|
||||
|
||||
@@ -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}
|
||||
)
|
||||
|
||||
@@ -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"]))
|
||||
|
||||
@@ -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 有限研究",
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
@@ -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):
|
||||
|
||||
@@ -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"))
|
||||
|
||||
Reference in New Issue
Block a user