diff --git a/.scratch/data-preparations/issues/01-implementation.md b/.scratch/data-preparations/issues/01-implementation.md new file mode 100644 index 0000000..c054e62 --- /dev/null +++ b/.scratch/data-preparations/issues/01-implementation.md @@ -0,0 +1,29 @@ +# 实现与验证 + +Type: task +Status: resolved + +按 spec.md 实现模型/接口、同步 CLI、页面及研究/AI/MCP 接入,更新测试和部署说明。 + +## Comments + +- 开始实施;工作区初始无改动。 + +- 完成模型与迁移 0015、集合/字段 API、完整同步 CLI、新页面与公共选择器、研究及 AI/MCP 输入快照接入;旧输入模型及接口已移除。 +- 全量同步复用账户事务锁实现同范围并发去重;重试也复用活动任务。逐集批次独立发布,保留分页检查点、失败原因和上一完整版本;连接故障暂停,等待超时不取消后台任务。 + +## Answer + +已完成本地实现与隔离验收,未提交或部署。 + +验证证据: + +- `uv run ruff check app tests`:通过;`git diff --check`:通过。 +- 后端全量 `uv run pytest -q`:420 项通过。之后的任务重试/CLI修改定向回归 62 项通过;补强模型不能追加准备集合的输入边界后,研究工作区 35 项通过。 +- `pnpm build`:通过,保留已有 bundle 大小与 lottie eval 警告。 +- Playwright:相关 15 项全部通过,覆盖数据目录、在线字段跨页多选、集合编辑、公共选择器、特征转换、模板/变体 AI、回测及来源快照、流水线、QuantFlow、导航和工作区;范围选项刷新调整后,再跑数据目录 2 项通过。已查看字段目录截图。 +- PostgreSQL 17 独立空库:0014 → 0015 迁移与 Alembic metadata check 通过;旧目录分页 offset 保留、旧输入表移除、分类筛选和分页、五路并发任务去重及冻结、删除集合后快照保留均通过。对应脚本 `backend/tests/preparations_postgres.py`,测试容器已清理。 +- CLI 隔离测试覆盖状态退出码、活动任务复用、超时、续页、网络暂停和重试;无效 delay/NaN 或零等待时间的实际进程退出码均为 2。 +- 生产说明已补充 1Panel 的 docker exec、日志、恢复命令、六小时默认等待和全部退出码;README 与 MCP 文档同步更新。 + +边界:全部上游为模拟数据,没有调用真实 WorldQuant,也没有配置或触发 1Panel 调度。真实平台过滤/分页协议、范围权限及调度效果需单独联调。迁移按用户确认不兼容旧研究输入,回退需要升级前数据库备份。 diff --git a/.scratch/data-preparations/spec.md b/.scratch/data-preparations/spec.md new file mode 100644 index 0000000..64dd320 --- /dev/null +++ b/.scratch/data-preparations/spec.md @@ -0,0 +1,11 @@ +# 数据目录、字段目录与数据准备重构 + +Status: ready-for-agent + +用户已批准实施。替换单数据集已保存输入,不兼容旧研究数据。数据准备是固定 instrument_type/region/universe/delay 的可编辑字段集合,可包含多个数据集;字段保留归属、描述、类型及来源。研究使用独立快照,集合编辑/删除不影响已创建研究。 + +数据目录行操作为查看、同步、使用;使用整集复制。字段目录提供 worldquant接口/本地同步两个 Tab、服务端丰富筛选与跨页多选,可新建或追加同范围集合。所有研究页面共用弹窗,并接入 AI/MCP。 + +全量同步仅通过 app.cli catalog-sync 创建持久化任务,由现有 runner 执行。支持范围、活动任务去重、逐集完整发布、检查点恢复、CLI 等待/退出码和 1Panel 调用说明。 + +验证:隔离 HTTP/数据库测试、后端 Ruff/pytest、前端构建与 Playwright、隔离 PostgreSQL 迁移。真实平台、部署和 1Panel 调度不在本地执行范围。 diff --git a/README.md b/README.md index dc8daab..f1de199 100644 --- a/README.md +++ b/README.md @@ -34,19 +34,23 @@ docker compose ps 工作空间和 AI 交互统一采用紧凑的 Lark 样式。Alpha 列表只滚动表体,分页保持在可用区域底部;个人信息页独立滚动。 -## 数据集与数据字段 +## 数据目录、字段目录与数据准备 -从侧栏进入“数据集”,设置 Region、Universe、Delay 后手动同步目录。范围选项表示本版支持的组合,平台账户实际权限以同步结果为准;分类和子分类来自已同步数据。 +统一流程是“查询或同步字段 → 整理数据准备集合 → 各研究模块选择集合”。 -选中一个数据集后默认使用整集字段;首次使用先同步全部字段。字段列表、搜索、类型、覆盖率、排序及翻页均不改变输入范围,只有明确取消勾选才排除字段。表头选择作用于整个已完成集合,支持恢复全选。字段与详情采用 75% / 30% 的工作区右抽屉,窄屏展开为全宽;逐层关闭保留父层条件。抽屉顶部可打开 AI 助手,业务抽屉暂时隐藏,收起助手后恢复;发送消息时会附带当前范围和输入引用;助手可通过工具读取本地目录和字段,不发送未保存的研究备注。 +数据目录行操作为“查看”“同步”“使用”。“同步目录”仅刷新数据集清单;行内“同步”更新全部字段。完整分页成功后才发布,失败、取消或刷新期间保留上一完整版本。“使用”复制当前完整字段到新准备集合,首次同步未完成时禁用。 -“用于 Alpha 模板”先保存输入草稿,在服务端固定数据集、研究范围、集合版本、字段 ID 和字段类型。点击“用此输入研究”将该快照带入聊天;也可从“已保存输入”恢复。后续同步不会改变旧草稿。 +字段目录分为 `worldquant接口` 和 `本地同步`。在线字段可直接加入集合,不标记数据集已同步。本地目录汇总已完整同步字段,提供范围、数据集、类型、关键词、分类、覆盖率、用户数、Alpha 数、同步时间和排序筛选。两种来源支持跨页勾选,表头选择只作用于当前页,切换范围清空选择。 -数据集和字段备注单独保存,版本冲突保留当前草稿。字段同步沿用已有任务面板的进度、取消、重试、等待连接和人工验证;每页与检查点同事务保存。只有完整分页成功才发布新集合,失败或取消继续使用上一版;首次未完成时不可准备输入。异常字段归属、覆盖率单位或分页协议会失败,不以部分字段代替全集。 +数据准备支持新建、修改名称和备注、复制、删除及批量删除,详情可查询、添加和移除字段并查看数据集归属。集合固定 `instrument_type + Region + Universe + Delay`,可包含多个同范围数据集;跨范围添加整批拒绝,重复字段去重。空集合可编辑但不能用于研究。 -增量迁移 `0003` 只增加目录、集合、备注与输入表,不改写旧迁移。`/api/v1/catalog` 提供带会话和来源校验的目录/字段查询、完整集合成员、备注、同步创建和输入草稿接口;创建目录同步返回任务 ID,查询、取消及重试仍使用 `/api/v1/sync-jobs`。新任务 `payload` 显式记录范围及数据集,保留旧 Alpha 任务契约。 +模板工坊、Alpha 变体、特征工程、回测研究、研究流水线和 QuantFlow 共用集合选择弹窗。选择保留集合 ID 和版本,提交研究时核对版本并固定完整字段快照;集合修改、删除或目录重同步均不改变已有研究。AI 可查询集合并固定输入;MCP 用 `search_data_preparations`、`get_data_preparation` 预览,以 `submit_backtests.preparation_refs` 提交按版本选择的集合。 -真实 WorldQuant 数据集 schema、字段所属数据集信息、0–1 覆盖率单位、范围权限和分页协议尚需只读联调。当前证据来自 HTTP 边界合成数据和隔离 PostgreSQL,不代表已验证真实平台兼容性。 +主要接口:`/api/v1/catalog/fields`(本地)、`/api/v1/catalog/worldquant/fields`(在线)、`/api/v1/data-preparations`(集合)和 `/api/v1/research/input-snapshots/{id}`(研究快照)。旧 `/catalog/inputs`、`/research/inputs` 和已保存输入入口已移除。迁移 `0015` 新建准备集合及独立快照表、移除旧输入表,不迁移旧研究数据;已有目录同步检查点保留。回退需要恢复升级前备份。 + +全量字段同步通过 `python -m app.cli catalog-sync` 入队,由现有单进程执行器处理。1Panel 夜间命令、日志、退出码和重试方法见 [生产部署说明](docs/deployment-gitea.md#5-1panel-夜间全量目录同步)。 + +实现验收使用模拟上游和隔离数据库。真实 WorldQuant 字段归属、过滤参数、分页协议、范围权限及 1Panel 调度效果需要单独联调。 ## AI 研究助手 diff --git a/backend/app/ai/contracts.py b/backend/app/ai/contracts.py index 19e89e5..9f1f112 100644 --- a/backend/app/ai/contracts.py +++ b/backend/app/ai/contracts.py @@ -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) diff --git a/backend/app/backtests/contracts.py b/backend/app/backtests/contracts.py index 4766a69..70bef83 100644 --- a/backend/app/backtests/contracts.py +++ b/backend/app/backtests/contracts.py @@ -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) diff --git a/backend/app/backtests/service.py b/backend/app/backtests/service.py index 7bdc379..78bce67 100644 --- a/backend/app/backtests/service.py +++ b/backend/app/backtests/service.py @@ -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"] = { diff --git a/backend/app/business.py b/backend/app/business.py index fbbd0e2..1396f51 100644 --- a/backend/app/business.py +++ b/backend/app/business.py @@ -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", diff --git a/backend/app/catalog/contracts.py b/backend/app/catalog/contracts.py index 52b61a8..280a9df 100644 --- a/backend/app/catalog/contracts.py +++ b/backend/app/catalog/contracts.py @@ -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] diff --git a/backend/app/catalog/routes.py b/backend/app/catalog/routes.py index 22e7eeb..c083274 100644 --- a/backend/app/catalog/routes.py +++ b/backend/app/catalog/routes.py @@ -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: diff --git a/backend/app/catalog/service.py b/backend/app/catalog/service.py index 9ac960a..8e61305 100644 --- a/backend/app/catalog/service.py +++ b/backend/app/catalog/service.py @@ -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) diff --git a/backend/app/catalog/sync.py b/backend/app/catalog/sync.py index 2964bc3..d044484 100644 --- a/backend/app/catalog/sync.py +++ b/backend/app/catalog/sync.py @@ -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 diff --git a/backend/app/cli.py b/backend/app/cli.py index 6e9ecc8..28f7643 100644 --- a/backend/app/cli.py +++ b/backend/app/cli.py @@ -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)) diff --git a/backend/app/jobs.py b/backend/app/jobs.py index 4a2757f..ba3c529 100644 --- a/backend/app/jobs.py +++ b/backend/app/jobs.py @@ -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, }, diff --git a/backend/app/main.py b/backend/app/main.py index 3cc2ce8..04cc184 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -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)) diff --git a/backend/app/mcp_api/server.py b/backend/app/mcp_api/server.py index 68d0ea1..3d43c56 100644 --- a/backend/app/mcp_api/server.py +++ b/backend/app/mcp_api/server.py @@ -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;缺缓存不自动刷新。"), diff --git a/backend/app/models.py b/backend/app/models.py index 9fed18b..d45b72a 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -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): diff --git a/backend/app/preparations/__init__.py b/backend/app/preparations/__init__.py new file mode 100644 index 0000000..af96426 --- /dev/null +++ b/backend/app/preparations/__init__.py @@ -0,0 +1 @@ +"""Data preparation collections and immutable research inputs.""" diff --git a/backend/app/preparations/contracts.py b/backend/app/preparations/contracts.py new file mode 100644 index 0000000..bd64db0 --- /dev/null +++ b/backend/app/preparations/contracts.py @@ -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 diff --git a/backend/app/preparations/routes.py b/backend/app/preparations/routes.py new file mode 100644 index 0000000..04094c7 --- /dev/null +++ b/backend/app/preparations/routes.py @@ -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}, + } diff --git a/backend/app/preparations/service.py b/backend/app/preparations/service.py new file mode 100644 index 0000000..69c0926 --- /dev/null +++ b/backend/app/preparations/service.py @@ -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 diff --git a/backend/app/research/ai_tools.py b/backend/app/research/ai_tools.py index 9bb7c83..f28fdea 100644 --- a/backend/app/research/ai_tools.py +++ b/backend/app/research/ai_tools.py @@ -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", diff --git a/backend/app/research/assets.py b/backend/app/research/assets.py index 66b0bfe..f83c77f 100644 --- a/backend/app/research/assets.py +++ b/backend/app/research/assets.py @@ -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 diff --git a/backend/app/research/contracts.py b/backend/app/research/contracts.py index 6f103b9..7826288 100644 --- a/backend/app/research/contracts.py +++ b/backend/app/research/contracts.py @@ -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") diff --git a/backend/app/research/experiments.py b/backend/app/research/experiments.py index 7431570..c4789a1 100644 --- a/backend/app/research/experiments.py +++ b/backend/app/research/experiments.py @@ -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]} diff --git a/backend/app/research/routes.py b/backend/app/research/routes.py index 983b5b2..393a5c0 100644 --- a/backend/app/research/routes.py +++ b/backend/app/research/routes.py @@ -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( diff --git a/backend/app/research/runtime.py b/backend/app/research/runtime.py index a27d3f0..1c6b882 100644 --- a/backend/app/research/runtime.py +++ b/backend/app/research/runtime.py @@ -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( diff --git a/backend/app/research/service.py b/backend/app/research/service.py index bc4fc23..83a2d7d 100644 --- a/backend/app/research/service.py +++ b/backend/app/research/service.py @@ -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, } ) diff --git a/backend/app/research/workflows.py b/backend/app/research/workflows.py index 67d88b5..57a5430 100644 --- a/backend/app/research/workflows.py +++ b/backend/app/research/workflows.py @@ -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 diff --git a/backend/app/research/workspace_contracts.py b/backend/app/research/workspace_contracts.py index 4ec2530..09a7723 100644 --- a/backend/app/research/workspace_contracts.py +++ b/backend/app/research/workspace_contracts.py @@ -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 diff --git a/backend/app/research_access/contracts.py b/backend/app/research_access/contracts.py index 34dee87..3a7d34a 100644 --- a/backend/app/research_access/contracts.py +++ b/backend/app/research_access/contracts.py @@ -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" diff --git a/backend/app/research_access/service.py b/backend/app/research_access/service.py index 08b1942..732b841 100644 --- a/backend/app/research_access/service.py +++ b/backend/app/research_access/service.py @@ -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"], diff --git a/backend/migrations/versions/0015_data_preparations.py b/backend/migrations/versions/0015_data_preparations.py new file mode 100644 index 0000000..76dc836 --- /dev/null +++ b/backend/migrations/versions/0015_data_preparations.py @@ -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("旧输入模型已移除;回退请恢复升级前数据库备份") diff --git a/backend/tests/catalog_migration_check.py b/backend/tests/catalog_migration_check.py index 6ce3b84..5854cc5 100644 --- a/backend/tests/catalog_migration_check.py +++ b/backend/tests/catalog_migration_check.py @@ -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" diff --git a/backend/tests/preparations_postgres.py b/backend/tests/preparations_postgres.py new file mode 100644 index 0000000..fa0be1b --- /dev/null +++ b/backend/tests/preparations_postgres.py @@ -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" + ) diff --git a/backend/tests/research_fake.py b/backend/tests/research_fake.py index 0251704..0e44894 100644 --- a/backend/tests/research_fake.py +++ b/backend/tests/research_fake.py @@ -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", diff --git a/backend/tests/research_flows_postgres.py b/backend/tests/research_flows_postgres.py index 1935ccb..8ff5b18 100644 --- a/backend/tests/research_flows_postgres.py +++ b/backend/tests/research_flows_postgres.py @@ -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 有限研究", diff --git a/backend/tests/research_live_acceptance.py b/backend/tests/research_live_acceptance.py index aacac37..840fdd5 100644 --- a/backend/tests/research_live_acceptance.py +++ b/backend/tests/research_live_acceptance.py @@ -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} ) diff --git a/backend/tests/research_outcomes_postgres.py b/backend/tests/research_outcomes_postgres.py index 24e8380..987896a 100644 --- a/backend/tests/research_outcomes_postgres.py +++ b/backend/tests/research_outcomes_postgres.py @@ -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"])) diff --git a/backend/tests/research_quantflow_postgres.py b/backend/tests/research_quantflow_postgres.py index 169f9e4..63577e9 100644 --- a/backend/tests/research_quantflow_postgres.py +++ b/backend/tests/research_quantflow_postgres.py @@ -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 有限研究", diff --git a/backend/tests/test_ai_capabilities.py b/backend/tests/test_ai_capabilities.py index 46bb9b5..5c25746 100644 --- a/backend/tests/test_ai_capabilities.py +++ b/backend/tests/test_ai_capabilities.py @@ -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" diff --git a/backend/tests/test_alpha_list_fields.py b/backend/tests/test_alpha_list_fields.py index 09b72b0..d8c3e3b 100644 --- a/backend/tests/test_alpha_list_fields.py +++ b/backend/tests/test_alpha_list_fields.py @@ -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() diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index 2e99395..3906943 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -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, } diff --git a/backend/tests/test_catalog.py b/backend/tests/test_catalog.py index 98037e7..84c82dd 100644 --- a/backend/tests/test_catalog.py +++ b/backend/tests/test_catalog.py @@ -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): diff --git a/backend/tests/test_mcp.py b/backend/tests/test_mcp.py index 6059311..d0da7e0 100644 --- a/backend/tests/test_mcp.py +++ b/backend/tests/test_mcp.py @@ -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"] diff --git a/backend/tests/test_preparations.py b/backend/tests/test_preparations.py new file mode 100644 index 0000000..f815537 --- /dev/null +++ b/backend/tests/test_preparations.py @@ -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 diff --git a/backend/tests/test_research_integration.py b/backend/tests/test_research_integration.py index 0b541d8..fdde5f0 100644 --- a/backend/tests/test_research_integration.py +++ b/backend/tests/test_research_integration.py @@ -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): diff --git a/backend/tests/test_research_workspace.py b/backend/tests/test_research_workspace.py index 1036ba6..fb54ea4 100644 --- a/backend/tests/test_research_workspace.py +++ b/backend/tests/test_research_workspace.py @@ -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")) diff --git a/docs/deployment-gitea.md b/docs/deployment-gitea.md index fcc44e2..2727759 100644 --- a/docs/deployment-gitea.md +++ b/docs/deployment-gitea.md @@ -87,3 +87,30 @@ docker logs --tail=100 wq-alpha-production-web-1 采用 compound 经验 1 中的生产/开发隔离、显式外部网络、必填凭据、固定项目名、独立迁移和部署健康检查。未照搬旧项目端口、业务 Job 或卷权限初始化:本应用后端不写 named volume,已有镜像使用 UID 10001。 凭据存放在 Gitea Secrets,通过步骤级环境变量交给 Compose,不落地到配置文件。Docker 容器仍需持有运行时配置,因此拥有 Runner 或 Docker 管理权限的人仍可能读取它们。若此前手动创建了旧 `.env.production`,新流程不再读取它;确认配置已迁移到 Gitea 并妥善备份密钥后可自行移除旧文件。相关官方资料:[Compose 外部网络](https://docs.docker.com/reference/compose-file/networks/)、[环境变量插值](https://docs.docker.com/compose/how-tos/environment-variables/variable-interpolation/)、[Gitea Runner 标签](https://gitea.com/gitea/runner/src/branch/main/README.md)。 + +## 5. 1Panel 夜间全量目录同步 + +在 1Panel 的计划任务中建立 Shell 脚本任务,执行周期由 1Panel 设置,例如每天深夜执行。每次命令明确指定一个范围: + +```bash +docker exec wq-alpha-production-backend-1 \ + python -m app.cli catalog-sync \ + --region USA --universe TOP3000 --delay 1 +``` + +`--instrument-type` 默认 `EQUITY`,`--wait-timeout` 默认 `21600` 秒(六小时)。容器已持有部署时注入的环境,定时脚本无需再次填写数据库、密码或加密密钥;不要加 `-it`。容器命令继承环境的行为见 [Docker exec 官方文档](https://docs.docker.com/reference/cli/docker/container/exec/)。启动前在网页完成 WorldQuant 连接,人工验证仍在网页处理。 + +CLI 只创建持久化任务并等待,不启动另一个同步执行器;后端必须正在运行。相同范围的活动全量任务复用同一任务 ID。任务先完整更新数据集清单,再逐个同步所有字段,每个数据集独立完整发布。普通单集失败保留上一版并继续其他数据集;鉴权和网络连接问题暂停全量任务。任务面板可查看阶段、当前数据集、字段分页位置、成功/失败统计,并取消或重试。 + +1Panel 执行日志保存 CLI 输出的任务 ID、状态变化、检查点和最终汇总。也可使用 `docker logs --tail=100 wq-alpha-production-backend-1` 排查执行器;命令不输出平台认证信息。不要公开包含研究数据的日志。 + +失败或部分失败时按日志中的任务 ID 从检查点继续,已成功发布的数据集不会重抓: + +```bash +docker exec wq-alpha-production-backend-1 \ + python -m app.cli catalog-sync --resume-job TASK_ID --wait-timeout 21600 +``` + +退出码:`0` 全部成功,`1` 失败/部分失败/取消,`2` 参数错误,`3` 需要连接或人工验证,`4` 等待超时。超时只结束 CLI 等待,后台任务继续;重复执行同范围命令可继续等待活动任务。恢复前先处理连接或验证问题。服务重启后由原执行器恢复持久化检查点;定时调度仅由 1Panel 负责。 + +真实平台协议和 1Panel 实际触发效果不属于模拟验收。 diff --git a/docs/mcp-research.md b/docs/mcp-research.md index a324883..36bcf1b 100644 --- a/docs/mcp-research.md +++ b/docs/mcp-research.md @@ -51,6 +51,8 @@ python -m app.cli mcp-token-revoke TOKEN_ID | authenticate_worldquant | `{action?:"connect"或"verify"}`;默认 connect,使用已保存凭据异步认证,返回 job_id;要求 research:refresh | | get_research_capabilities | `{}`,含单次候选上限及完整候选 schema | | create_research_template | `{template,hypothesis,source_item_ids,idempotency_key,reference?}`;保存调用方生成的完整模板,要求 research:write | +| search_data_preparations | `{q?,scope_key?,limit?,offset?}`;查询可编辑集合及当前版本 | +| get_data_preparation | `{id,version,q?,limit?,offset?}`;按版本分页预览字段与数据集归属,版本冲突重新选择 | | search_catalog | `{filters:{region,universe,delay,...},dataset_id?}`;省略 dataset_id 查数据集,提供则查字段 | | get_research_metadata | `{query:{kind,...}}`;kind 为 scopes/settings/operators/field_availability | | refresh_research_data | `{query:{kind,...}}`;kind 为 catalog/operators/settings/field_availability/pnl | @@ -60,7 +62,7 @@ python -m app.cli mcp-token-revoke TOKEN_ID | check_self_correlation | `{alpha_ids:[...]}`,1–100 个已导入 Alpha ID;异步返回 job_id | | get_self_correlation | `{alpha_id}`;只读最新本地结果,含缓存和 stale 状态 | | search_backtests | 来源、reference、status、带时区起止时间、scope、q、候选精确匹配及分页 | -| submit_backtests | `{name,candidates,idempotency_key,duplicate_policy?,source?}` | +| submit_backtests | `{name,candidates,idempotency_key,preparation_refs?,duplicate_policy?,source?}` | | get_backtest | `{run_id,after?,event_limit?}`,after 为事件游标 | | get_backtest_results | `{run_id,item_ids?,limit?,offset?}` | | get_backtest_artifact | `{item_id,kind,limit?,offset?,date_from?,date_to?}`,kind 为 snapshot/pnl | @@ -188,3 +190,5 @@ MCP_TEST_DATABASE_URL=postgresql+asyncpg://USER:PASSWORD@127.0.0.1:PORT/wq_mcp_t ``` 该脚本执行迁移、并发提交/控制、重启重放、回退及重升级,不用于个人库或生产库。生产启用、真实平台兼容性、真实额度和客户端实际凭据配置仍需另行授权验证。本功能不会恢复任何定时研究。 + +使用数据准备集合时,`preparation_refs` 为最多 20 个 `{id,version}`。提交时核对版本、范围和字段,固定独立快照并保存到回测来源;空集合或版本冲突不会创建运行。后续编辑或删除集合不影响回测。无需旧输入草稿接口。 diff --git a/frontend/src/App.tsx b/frontend/src/App.tsx index 2de49b2..1f90ee1 100644 --- a/frontend/src/App.tsx +++ b/frontend/src/App.tsx @@ -1,3 +1,6 @@ +import { FieldDirectory } from "./preparations/FieldDirectory"; +import { DataPreparationPage } from "./preparations/DataPreparationPage"; +import { SnapshotDialog } from "./preparations/SnapshotDialog"; import { useCallback, useEffect, useRef, useState } from "react"; import { Badge, @@ -390,6 +393,15 @@ export default function App() { onTask={taskCreated} /> +