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

This commit is contained in:
yuxuanhui
2026-09-12 01:24:02 +08:00
parent 849f86fef7
commit 394438e753
82 changed files with 4146 additions and 2076 deletions
@@ -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 调度。真实平台过滤/分页协议、范围权限及调度效果需单独联调。迁移按用户确认不兼容旧研究输入,回退需要升级前数据库备份。
+11
View File
@@ -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 调度不在本地执行范围。
+11 -7
View File
@@ -34,19 +34,23 @@ docker compose ps
工作空间和 AI 交互统一采用紧凑的 Lark 样式。Alpha 列表只滚动表体,分页保持在可用区域底部;个人信息页独立滚动。 工作空间和 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 研究助手 ## AI 研究助手
+3 -1
View File
@@ -47,6 +47,8 @@ class PageContext(Contract):
"alphas", "alphas",
"account", "account",
"datasets", "datasets",
"fields",
"preparations",
"backtests", "backtests",
"operators", "operators",
"templates", "templates",
@@ -62,7 +64,7 @@ class PageContext(Contract):
dataset_id: str | None = Field(default=None, min_length=1, max_length=200) 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) 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) 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 unsaved_field_selection: bool = False
backtest_run_id: str | None = Field(default=None, max_length=36) backtest_run_id: str | None = Field(default=None, max_length=36)
backtest_preview_id: str | None = Field(default=None, max_length=36) backtest_preview_id: str | None = Field(default=None, max_length=36)
+7 -1
View File
@@ -6,6 +6,7 @@ from typing import Literal
from pydantic import Field, field_validator, model_validator from pydantic import Field, field_validator, model_validator
from ..preparations.contracts import PreparationReference
from ..schemas import Contract from ..schemas import Contract
@@ -48,13 +49,18 @@ class Source(Contract):
kind: str = Field(default="manual", min_length=1, max_length=100) kind: str = Field(default="manual", min_length=1, max_length=100)
reference: str | None = Field(default=None, max_length=200) reference: str | None = Field(default=None, max_length=200)
batch_id: 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) research_id: str | None = Field(default=None, max_length=200)
parent_run_id: str | None = Field(default=None, max_length=36) parent_run_id: str | None = Field(default=None, max_length=36)
hypothesis: str | None = Field(default=None, max_length=2000) hypothesis: str | None = Field(default=None, max_length=2000)
class DraftInput(Contract): 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) name: str = Field(min_length=1, max_length=200)
source: Source = Field(default_factory=Source) source: Source = Field(default_factory=Source)
candidates: list[Candidate] = Field(min_length=1, max_length=10000) candidates: list[Candidate] = Field(min_length=1, max_length=10000)
+27 -1
View File
@@ -106,8 +106,33 @@ class Backtests:
"limits": "并发和批大小为本地配置,并非平台可用额度;账户外提交不在预算内", "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): 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: if draft_id:
changed = await self.db.execute( changed = await self.db.execute(
update(BacktestDraft) update(BacktestDraft)
@@ -168,6 +193,7 @@ class Backtests:
producer. ai_context separately identifies whoever starts the execution. producer. ai_context separately identifies whoever starts the execution.
""" """
if body.inline: if body.inline:
await self.bind_preparations(body.inline)
data = body.inline.model_dump(mode="json") data = body.inline.model_dump(mode="json")
if self.ai_context and not preserve_source: if self.ai_context and not preserve_source:
data["source"] = { data["source"] = {
+9 -1
View File
@@ -263,9 +263,17 @@ class Business:
return {"ok": True, "job_id": job_id} return {"ok": True, "job_id": job_id}
async def retry_job(self, 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: if not job:
raise HTTPException(404, "任务不存在") 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 ( if job.status not in (
"failed", "failed",
"cancelled", "cancelled",
+1 -27
View File
@@ -3,7 +3,7 @@
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Annotated, Literal from typing import Annotated, Literal
from pydantic import AfterValidator, BaseModel, Field, model_validator from pydantic import AfterValidator, BaseModel, Field
from ..schemas import Contract from ..schemas import Contract
@@ -54,20 +54,6 @@ class NoteInput(Contract):
version: int = Field(ge=1) 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): class NoteOutput(BaseModel):
note: str note: str
version: int version: int
@@ -107,18 +93,6 @@ class CatalogPage(BaseModel):
field_types: list[str] = Field(default_factory=list) 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): class CollectionOutput(BaseModel):
collection_version: str | None collection_version: str | None
field_ids: list[str] field_ids: list[str]
-20
View File
@@ -12,8 +12,6 @@ from .contracts import (
CatalogPage, CatalogPage,
CollectionOutput, CollectionOutput,
EntryOutput, EntryOutput,
InputOutput,
InputPreparation,
NoteInput, NoteInput,
NoteOutput, NoteOutput,
Scope, Scope,
@@ -76,24 +74,6 @@ async def sync(request: Request, body: CatalogJobInput):
return result 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) @router.get("/datasets/{dataset_id}/collection", response_model=CollectionOutput)
async def collection(request: Request, dataset_id: str, scope: Annotated[Scope, Query()]): async def collection(request: Request, dataset_id: str, scope: Annotated[Scope, Query()]):
async with request.app.state.sessions() as db: async with request.app.state.sessions() as db:
+5 -65
View File
@@ -4,7 +4,6 @@ The dataset row serializes collection publication and draft creation on PostgreS
No page filters participate in template input selection. No page filters participate in template input selection.
""" """
from datetime import timezone
from uuid import uuid4 from uuid import uuid4
from fastapi import HTTPException from fastapi import HTTPException
@@ -18,7 +17,6 @@ from ..models import (
CatalogNote, CatalogNote,
CatalogScope, CatalogScope,
Job, Job,
TemplateInput,
now, now,
) )
from ..schemas import JobOutput from ..schemas import JobOutput
@@ -164,13 +162,13 @@ class Catalog:
raise HTTPException(409, "研究备注已被修改;当前草稿已保留,请载入最新记录后重新保存") raise HTTPException(409, "研究备注已被修改;当前草稿已保留,请载入最新记录后重新保存")
return dict(note=body.note, version=body.version + 1, updated_at=now()) 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()) 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"): if not account.password_encrypted or account.connection_status in ("disconnected", "error"):
raise HTTPException(409, "请先连接 WorldQuant") raise HTTPException(409, "请先连接 WorldQuant")
if body.dataset_id: if body.dataset_id:
await self.dataset(body.scope, 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") payload = body.model_dump(mode="json")
jobs = ( jobs = (
await self.db.scalars( await self.db.scalars(
@@ -190,7 +188,7 @@ class Catalog:
job = Job(id=str(uuid4()), kind=kind, payload=payload) job = Job(id=str(uuid4()), kind=kind, payload=payload)
self.db.add(job) self.db.add(job)
await self.db.flush() 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() await self.db.flush()
return JobOutput.model_validate(job) return JobOutput.model_validate(job)
@@ -210,64 +208,6 @@ class Catalog:
) )
return dict(collection_version=dataset.field_version, field_ids=ids) 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): async def input(self, input_id):
row = await self.db.get(TemplateInput, input_id) from ..preparations.service import Preparations
if not row: return await Preparations(self.db).snapshot(input_id)
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]
+82 -14
View File
@@ -4,6 +4,7 @@ import asyncio
import math import math
import re import re
from urllib.parse import parse_qs, urlparse from urllib.parse import parse_qs, urlparse
from uuid import uuid4
from sqlalchemy import select 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"]) scope = Scope.model_validate(payload["scope"])
dataset_id = payload.get("dataset_id") dataset_id = payload.get("dataset_id")
batch_id = batch_id or job_id
async with runner.sessions() as db: async with runner.sessions() as db:
checkpoint = (await db.get(Job, job_id)).checkpoint batch = await db.get(CatalogBatch, batch_id)
if checkpoint.get("done"): if batch.complete:
return return
offset = checkpoint.get("offset", 0) offset = batch.offset
while True: while True:
await runner.checkpoint(job_id, {"next_retry_at": None}) await runner.checkpoint(job_id, {"next_retry_at": None})
raw = await runner.client.catalog_page(scope.model_dump(), dataset_id, offset) 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) job = await db.get(Job, job_id)
if job.cancel_requested: if job.cancel_requested:
raise asyncio.CancelledError() raise asyncio.CancelledError()
batch = await db.get(CatalogBatch, job_id) batch = await db.get(CatalogBatch, batch_id)
added = 0 added = 0
for entry in entries: for entry in entries:
if await db.get(CatalogEntry, (job_id, entry["id"])): if await db.get(CatalogEntry, (batch_id, entry["id"])):
continue continue
db.add(CatalogEntry(batch_id=job_id, **entry)) db.add(CatalogEntry(batch_id=batch_id, **entry))
await db.flush() await db.flush()
added += 1 added += 1
owner = dataset_id or entry["id"] owner = dataset_id or entry["id"]
@@ -112,25 +114,28 @@ async def sync_catalog(runner, job_id, payload):
if rows and not added: if rows and not added:
raise WqError("平台分页重复且未前进,已保留进度", "invalid_response") raise WqError("平台分页重复且未前进,已保留进度", "invalid_response")
batch.count += added batch.count += added
job.processed = batch.count if not full:
job.processed = batch.count
offset += len(rows) 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() job.updated_at = now()
if not more: if not more:
batch.complete, batch.completed_at = True, now() batch.complete, batch.completed_at = True, now()
job.total = batch.count if not full:
job.total = batch.count
if dataset_id: if dataset_id:
dataset = await db.scalar( dataset = await db.scalar(
select(CatalogDataset) select(CatalogDataset)
.where(CatalogDataset.scope_key == scope.key(), CatalogDataset.id == dataset_id) .where(CatalogDataset.scope_key == scope.key(), CatalogDataset.id == dataset_id)
.with_for_update() .with_for_update()
) )
dataset.field_version = job_id dataset.field_version = batch_id
else: else:
scope_row = await db.get(CatalogScope, scope.key()) 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 = ( 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() ).all()
for item_id in ids: for item_id in ids:
if not await db.get(CatalogDataset, (scope.key(), item_id)): 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() await db.commit()
if not more: if not more:
return return
async def sync_full_catalog(runner, job_id, payload):
"""Resume each dataset batch independently; only publish complete enumerations."""
scope = Scope.model_validate(payload["scope"])
options = await runner.client.get_platform_setting_options()
if not any(r["instrument_type"] == scope.instrument_type and r["region"] == scope.region
and r["delay"] == scope.delay and scope.universe in r["universes"]
for r in options["instrument_options"]):
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.checkpoint = {**job.checkpoint, "error_code": "invalid_scope"}
raise WqError("平台不支持该研究范围", "invalid_scope")
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.checkpoint = {**{k: v for k, v in job.checkpoint.items() if k != "error_code"}, "phase": "catalog"}
await sync_catalog(runner, job_id, {"scope": payload["scope"]}, full=True)
async with runner.sessions() as db:
ids = list(await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id)
.order_by(CatalogEntry.id)))
job = await db.get(Job, job_id)
failures = dict(job.checkpoint.get("failures", {}))
completed = 0
await runner.checkpoint(job_id, {"total": len(ids)})
for dataset_id in ids:
async with runner.sessions.begin() as db:
batch = await db.scalar(select(CatalogBatch).where(CatalogBatch.job_id == job_id,
CatalogBatch.dataset_id == dataset_id))
if not batch:
batch = CatalogBatch(id=str(uuid4()), job_id=job_id, scope_key=scope.key(), dataset_id=dataset_id)
db.add(batch)
await db.flush()
batch_id, complete = batch.id, batch.complete
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
job.checkpoint = {**job.checkpoint, "phase": "fields", "dataset_id": dataset_id,
"datasets_completed": completed, "datasets_total": len(ids),
"offset": batch.offset, "current_field_count": batch.count}
if not complete:
try:
await sync_catalog(runner, job_id, {"scope": payload["scope"], "dataset_id": dataset_id},
batch_id=batch_id, full=True)
except WqError as exc:
if exc.code in ("disconnected", "authentication_failed", "identity_mismatch", "verification_required", "network_error"):
raise
failures[dataset_id] = str(exc)
if complete or (await _batch_complete(runner, batch_id)):
failures.pop(dataset_id, None)
completed += 1
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.processed, job.failed = completed, len(failures)
job.checkpoint = {**job.checkpoint, "datasets_completed": completed, "failures": failures}
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.error = f"{len(failures)} 个数据集同步失败" if failures else None
job.checkpoint = {**job.checkpoint, "phase": "finished"}
async def _batch_complete(runner, batch_id):
async with runner.sessions() as db:
return (await db.get(CatalogBatch, batch_id)).complete
+69
View File
@@ -3,6 +3,7 @@
import argparse import argparse
import asyncio import asyncio
import getpass import getpass
import math
from sqlalchemy import delete, update from sqlalchemy import delete, update
@@ -59,6 +60,63 @@ async def token_command(args):
await engine.dispose() 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__": if __name__ == "__main__":
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
commands = parser.add_subparsers(dest="command", required=True) commands = parser.add_subparsers(dest="command", required=True)
@@ -70,8 +128,19 @@ if __name__ == "__main__":
commands.add_parser("mcp-token-list") commands.add_parser("mcp-token-list")
revoke = commands.add_parser("mcp-token-revoke") revoke = commands.add_parser("mcp-token-revoke")
revoke.add_argument("token_id") 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() 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: 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)) asyncio.run(reset_password() if args.command == "reset-password" else token_command(args))
except ValueError as exc: except ValueError as exc:
parser.error(str(exc)) parser.error(str(exc))
+5 -1
View File
@@ -255,6 +255,10 @@ class Runner:
await self.ensure_connected(force=kind == "connect") await self.ensure_connected(force=kind == "connect")
if kind in ("connect", "profile"): if kind in ("connect", "profile"):
await self.refresh_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"): elif kind in ("catalog_sync", "field_sync"):
from .catalog.sync import sync_catalog from .catalog.sync import sync_catalog
@@ -298,7 +302,7 @@ class Runner:
await self.checkpoint( await self.checkpoint(
job_id, 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), "error": str(exc),
"next_retry_at": None, "next_retry_at": None,
}, },
+2
View File
@@ -27,6 +27,7 @@ from .db import create_database
from .jobs import AUTH_KINDS, Runner, create_job from .jobs import AUTH_KINDS, Runner, create_job
from .mcp_api.token_routes import router as mcp_token_router from .mcp_api.token_routes import router as mcp_token_router
from .models import Account, Admin, BacktestConfig, Job, JobItem, LoginSession 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.routes import router as research_router
from .research.runtime import ResearchRuntime from .research.runtime import ResearchRuntime
from .schemas import ( 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(backtest_router)
app.include_router(api) app.include_router(api)
app.include_router(catalog_router) app.include_router(catalog_router)
app.include_router(preparations_router)
app.include_router(research_catalog_router) app.include_router(research_catalog_router)
app.include_router(research_router) app.include_router(research_router)
app.include_router(ai_router(ai_runtime)) app.include_router(ai_router(ai_runtime))
+3 -1
View File
@@ -21,6 +21,8 @@ from ..research_access.service import ResearchAccess, ResearchError
# Name, schema, business method, required scope, description. No generic arbitrary HTTP tool. # Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
TOOLS = { 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、版本和理论组合数;仅核验结构及来源,不验证所有参数组合,不再次调用模型、不执行回测、不覆盖已有模板。"), "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。"), "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。"), "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 读结果。"), "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 状态;无缓存不自动检查,结果不等同于平台提交资格。"), "get_self_correlation": (c.SelfCorrelationReference, "self_correlation", "research:read", "读取指定 Alpha 最新的本地自相关缓存,包括最大相关系数、样本覆盖与 stale 状态;无缓存不自动检查,结果不等同于平台提交资格。"),
"search_backtests": (c.History, "history", "research:read", "分页查历史候选与固定设置;candidates 按完整输入精确匹配,不推断数学等价。"), "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": (c.RunReference, "run", "research:read", "读取真实运行进度、提交数量和可选增量事件;受理不等于成功。"),
"get_backtest_results": (c.Results, "results", "research:read", "分页读取固定快照指标、全部非通过检查及三层状态;缺失指标不补零。"), "get_backtest_results": (c.Results, "results", "research:read", "分页读取固定快照指标、全部非通过检查及三层状态;缺失指标不补零。"),
"get_backtest_artifact": (c.Artifact, "artifact", "research:read", "分页读取候选脱敏快照的顶层键值或独立采集的 PnL;缺缓存不自动刷新。"), "get_backtest_artifact": (c.Artifact, "artifact", "research:read", "分页读取候选脱敏快照的顶层键值或独立采集的 PnL;缺缓存不自动刷新。"),
+31 -9
View File
@@ -340,7 +340,9 @@ class CatalogScope(Base):
class CatalogBatch(Base): class CatalogBatch(Base):
__tablename__ = "catalog_batches" __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) scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True)
dataset_id: Mapped[str | None] = mapped_column(String(200)) dataset_id: Mapped[str | None] = mapped_column(String(200))
complete: Mapped[bool] = mapped_column(Boolean, default=False) 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) updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
class TemplateInput(Base): class DataPreparation(Base):
__tablename__ = "template_inputs" """Editable collection; scope never changes after creation."""
__tablename__ = "data_preparations"
id: Mapped[str] = mapped_column(String(36), primary_key=True) id: Mapped[str] = mapped_column(String(36), primary_key=True)
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True) name: Mapped[str] = mapped_column(String(200))
dataset_id: Mapped[str] = mapped_column(String(200)) note: Mapped[str] = mapped_column(Text, default="")
collection_version: Mapped[str] = mapped_column(ForeignKey("catalog_batches.id")) scope_key: Mapped[str] = mapped_column(String(200), index=True)
selection: Mapped[str] = mapped_column(String(20)) scope: Mapped[dict] = mapped_column(JSON)
field_ids: Mapped[list] = mapped_column(JSON) version: Mapped[int] = mapped_column(Integer, default=1)
field_types: Mapped[dict] = mapped_column(JSON)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now) 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): class CatalogResource(Base):
+1
View File
@@ -0,0 +1 @@
"""Data preparation collections and immutable research inputs."""
+80
View File
@@ -0,0 +1,80 @@
"""Shared collection and field-query contracts."""
from datetime import datetime
from typing import Literal
from pydantic import Field, model_validator
from ..catalog.contracts import Scope
from ..schemas import Contract
class FieldFilters(Scope):
q: str = Field(default="", max_length=300)
dataset_id: str | None = None
category: str | None = None
subcategory: str | None = None
field_type: str | None = None
coverage_min: float | None = Field(default=None, ge=0, le=1)
coverage_max: float | None = Field(default=None, ge=0, le=1)
user_count_min: int | None = Field(default=None, ge=0)
user_count_max: int | None = Field(default=None, ge=0)
alpha_count_min: int | None = Field(default=None, ge=0)
alpha_count_max: int | None = Field(default=None, ge=0)
synced_from: datetime | None = None
synced_to: datetime | None = None
sort: Literal["id", "name", "dataset_id", "coverage", "user_count", "alpha_count", "synced_at"] = "id"
direction: Literal["asc", "desc"] = "asc"
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
@model_validator(mode="after")
def ranges(self):
for key in ("coverage", "user_count", "alpha_count"):
low, high = getattr(self, key + "_min"), getattr(self, key + "_max")
if low is not None and high is not None and low > high:
raise ValueError("筛选下限不能超过上限")
return self
class FieldReference(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
field_id: str = Field(min_length=1, max_length=200)
source: Literal["local", "worldquant"] = "local"
collection_version: str | None = None
class PreparationCreate(Contract):
name: str = Field(min_length=1, max_length=200)
note: str = Field(default="", max_length=20000)
scope: Scope
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
class PreparationVersion(Contract):
version: int = Field(ge=1)
class PreparationEdit(PreparationVersion):
name: str = Field(min_length=1, max_length=200)
note: str = Field(default="", max_length=20000)
class MemberChange(PreparationVersion):
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
remove_ids: list[str] = Field(default_factory=list, max_length=10000)
class PreparationReference(PreparationVersion):
id: str = Field(min_length=1, max_length=36)
class PreparationReferences(Contract):
items: list[PreparationReference] = Field(min_length=1, max_length=100)
class DatasetCopy(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
collection_version: str
+195
View File
@@ -0,0 +1,195 @@
"""Authenticated preparation and field-directory endpoints."""
from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import delete, select
from ..catalog.contracts import Scope
from ..models import CatalogScope, PreparationField, now
from ..security import require_auth
from .contracts import (
DatasetCopy,
FieldFilters,
MemberChange,
PreparationCreate,
PreparationEdit,
PreparationReferences,
PreparationVersion,
)
from .service import Preparations
router = APIRouter(prefix="/api/v1", dependencies=[Depends(require_auth)], tags=["data-preparations"])
@router.get("/catalog/local-scopes")
async def local_scopes(request: Request):
async with request.app.state.sessions() as db:
rows = await db.scalars(select(CatalogScope))
options = {}
for row in rows:
key = (row.scope["instrument_type"], row.scope["region"], row.scope["delay"])
option = options.setdefault(
key, {k: row.scope[k] for k in ("instrument_type", "region", "delay")}
)
option.setdefault("universes", []).append(row.scope["universe"])
return {"instrument_options": list(options.values())}
@router.get("/catalog/fields")
async def local_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
async with request.app.state.sessions() as db:
return await Preparations(db).fields(filters)
@router.get("/catalog/worldquant/fields")
async def online_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
async with request.app.state.sessions() as db:
return await Preparations(db, request.app.state.runner.client).online_fields(filters)
@router.get("/data-preparations")
async def preparations(
request: Request,
q: str = "",
scope_key: str | None = None,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
):
async with request.app.state.sessions() as db:
return await Preparations(db).list(q, scope_key, limit, offset)
@router.post("/data-preparations", status_code=201)
async def create(request: Request, body: PreparationCreate):
async with request.app.state.sessions.begin() as db:
service = Preparations(db, request.app.state.runner.client)
fields = await service.resolve_fields(body.scope, body.fields)
return await service.create(body.name, body.note, body.scope, fields)
@router.post("/data-preparations/from-dataset", status_code=201)
async def from_dataset(request: Request, body: DatasetCopy):
async with request.app.state.sessions.begin() as db:
return await Preparations(db).copy_dataset(body)
@router.post("/data-preparations/batch-delete")
async def batch_delete(request: Request, body: PreparationReferences):
async with request.app.state.sessions.begin() as db:
return await Preparations(db).remove(body.items)
@router.post("/data-preparations/freeze", status_code=201)
async def freeze(request: Request, body: PreparationReferences):
async with request.app.state.sessions.begin() as db:
return {"items": await Preparations(db).freeze(body.items)}
@router.get("/research/input-snapshots/{snapshot_id}")
async def snapshot(request: Request, snapshot_id: str):
async with request.app.state.sessions() as db:
return await Preparations(db).snapshot(snapshot_id)
@router.get("/data-preparations/{preparation_id}")
async def detail(request: Request, preparation_id: str):
async with request.app.state.sessions() as db:
service = Preparations(db)
return await service.output(await service.get(preparation_id))
@router.patch("/data-preparations/{preparation_id}")
async def edit(request: Request, preparation_id: str, body: PreparationEdit):
async with request.app.state.sessions.begin() as db:
service = Preparations(db)
row = await service.get(preparation_id, body.version, lock=True)
row.name, row.note, row.updated_at, row.version = body.name, body.note, now(), row.version + 1
return await service.output(row)
@router.delete("/data-preparations/{preparation_id}")
async def remove(request: Request, preparation_id: str, version: int = Query(ge=1)):
from .contracts import PreparationReference
async with request.app.state.sessions.begin() as db:
return await Preparations(db).remove([PreparationReference(id=preparation_id, version=version)])
@router.post("/data-preparations/{preparation_id}/copy", status_code=201)
async def copy(request: Request, preparation_id: str, body: PreparationVersion):
async with request.app.state.sessions.begin() as db:
service = Preparations(db)
row = await service.get(preparation_id, body.version, lock=True)
fields = [
f.content
for f in await db.scalars(
select(PreparationField).where(PreparationField.preparation_id == row.id)
)
]
return await service.create(
(row.name + " 副本")[:200], row.note, Scope.model_validate(row.scope), fields
)
@router.get("/data-preparations/{preparation_id}/fields")
async def members(
request: Request,
preparation_id: str,
q: str = "",
dataset_id: str | None = None,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
):
async with request.app.state.sessions() as db:
return await Preparations(db).members(preparation_id, q, dataset_id, limit, offset)
@router.patch("/data-preparations/{preparation_id}/fields")
async def change_members(request: Request, preparation_id: str, body: MemberChange):
async with request.app.state.sessions.begin() as db:
service = Preparations(db, request.app.state.runner.client)
row = await service.get(preparation_id, body.version, lock=True)
fields = await service.resolve_fields(Scope.model_validate(row.scope), body.fields)
await service.add(row, fields)
if body.remove_ids:
present = set(
await db.scalars(
select(PreparationField.field_id).where(
PreparationField.preparation_id == row.id,
PreparationField.field_id.in_(body.remove_ids),
)
)
)
if set(body.remove_ids) - present:
raise HTTPException(422, "移除项含不属于该集合的字段")
await db.execute(
delete(PreparationField).where(
PreparationField.preparation_id == row.id, PreparationField.field_id.in_(body.remove_ids)
)
)
row.version, row.updated_at = row.version + 1, now()
return await service.output(row)
@router.get("/data-preparations/{preparation_id}/selection")
async def selection(request: Request, preparation_id: str, version: int = Query(ge=1)):
async with request.app.state.sessions() as db:
service = Preparations(db)
row = await service.get(preparation_id, version)
fields = [
f.content
for f in await db.scalars(
select(PreparationField)
.where(PreparationField.preparation_id == row.id)
.order_by(PreparationField.field_id)
)
]
return {
**await service.output(row),
"fields": fields,
"field_ids": [f["id"] for f in fields],
"field_types": {f["id"]: f["field_type"] for f in fields},
"dataset_ids": sorted({f["dataset_id"] for f in fields}),
"preparation_ref": {"id": row.id, "version": row.version},
}
+450
View File
@@ -0,0 +1,450 @@
"""Collection operations own validation; callers own transactions and authorization."""
from uuid import uuid4
from fastapi import HTTPException
from sqlalchemy import delete, func, or_, select
from sqlalchemy.orm import aliased
from ..catalog.contracts import EntryOutput, Scope
from ..catalog.research_metadata import upstream
from ..catalog.service import Catalog
from ..catalog.sync import identifier, label, normalize
from ..models import (
CatalogBatch,
CatalogDataset,
CatalogEntry,
CatalogScope,
DataPreparation,
PreparationField,
ResearchInputSnapshot,
now,
)
from ..research.serialization import encode_snapshot
from .contracts import FieldFilters
def page(items, total, limit, offset, **extra):
return dict(
items=items, total=total, limit=limit, offset=offset, has_more=offset + len(items) < total, **extra
)
def contains(value):
return "%" + value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
class Preparations:
def __init__(self, db, client=None):
self.db, self.client = db, client
async def fields(self, filters):
scope = await self.db.get(CatalogScope, filters.key())
owner = aliased(CatalogEntry)
query = (
select(
CatalogEntry,
CatalogBatch.dataset_id,
owner.name.label("dataset_name"),
owner.category,
owner.subcategory,
)
.join(CatalogBatch, CatalogEntry.batch_id == CatalogBatch.id)
.join(
CatalogDataset,
(CatalogDataset.field_version == CatalogBatch.id)
& (CatalogDataset.scope_key == filters.key()),
)
.outerjoin(
owner,
(owner.id == CatalogBatch.dataset_id)
& (owner.batch_id == (scope.catalog_version if scope else None)),
)
.where(CatalogBatch.complete.is_(True))
)
if filters.q:
query = query.where(
or_(
*[
column.ilike(contains(filters.q), escape="\\")
for column in (
CatalogEntry.id,
CatalogEntry.name,
CatalogEntry.description,
CatalogBatch.dataset_id,
owner.name,
)
]
)
)
for key in ("dataset_id", "field_type", "category", "subcategory"):
value = getattr(filters, key)
column = (
CatalogBatch.dataset_id
if key == "dataset_id"
else func.coalesce(getattr(CatalogEntry, key), getattr(owner, key))
if key in ("category", "subcategory")
else getattr(CatalogEntry, key)
)
if value:
query = query.where(column == value)
for key in ("coverage", "user_count", "alpha_count"):
low, high = getattr(filters, key + "_min"), getattr(filters, key + "_max")
if low is not None:
query = query.where(getattr(CatalogEntry, key) >= low)
if high is not None:
query = query.where(getattr(CatalogEntry, key) <= high)
if filters.synced_from:
query = query.where(CatalogEntry.synced_at >= filters.synced_from)
if filters.synced_to:
query = query.where(CatalogEntry.synced_at <= filters.synced_to)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
column = (
CatalogBatch.dataset_id if filters.sort == "dataset_id" else getattr(CatalogEntry, filters.sort)
)
query = query.order_by(
(column.desc() if filters.direction == "desc" else column.asc()).nulls_last(),
CatalogBatch.dataset_id,
CatalogEntry.id,
)
rows = (await self.db.execute(query.limit(filters.limit).offset(filters.offset))).all()
items = [
{
**self.local_field(row, dataset_id, name, filters),
"category": row.category or category,
"subcategory": row.subcategory or subcategory,
}
for row, dataset_id, name, category, subcategory in rows
]
return page(items, total, filters.limit, filters.offset)
@staticmethod
def local_field(row, dataset_id, name, scope):
return encode_snapshot(
dict(
**EntryOutput.model_validate(row, from_attributes=True).model_dump(
exclude={"scope", "dataset_id", "collection_version"}
),
field_id=row.id,
dataset_id=dataset_id,
dataset_name=name or dataset_id,
collection_version=row.batch_id,
scope=scope.model_dump(include=set(Scope.model_fields)),
source="local",
fetched_at=row.synced_at,
)
)
async def online_fields(self, filters):
if self.client is None:
raise HTTPException(409, "请先连接 WorldQuant")
params = dict(
instrumentType=filters.instrument_type,
region=filters.region,
universe=filters.universe,
delay=filters.delay,
limit=filters.limit,
offset=filters.offset,
)
for key, remote in (("q", "search"), ("dataset_id", "dataset.id"), ("field_type", "type")):
value = getattr(filters, key)
if value:
params[remote] = value
for key, remote in (
("coverage", "coverage"),
("user_count", "userCount"),
("alpha_count", "alphaCount"),
):
for suffix, op in (("min", ">"), ("max", "<")):
value = getattr(filters, key + "_" + suffix)
if value is not None:
params[remote + op] = value
raw = await upstream(self.client.get("/data-fields", params))
rows = raw.get("results")
if not isinstance(rows, list):
raise HTTPException(502, "平台字段列表格式无法识别")
items = []
for row in rows:
if not isinstance(row, dict):
raise HTTPException(502, "平台字段记录格式无法识别")
dataset = row.get("dataset")
owner = dataset.get("id") if isinstance(dataset, dict) else dataset
try:
owner = identifier(owner)
data = normalize(row, owner)
except Exception as exc:
from ..worldquant import WqError
if isinstance(exc, WqError):
raise HTTPException(502, str(exc)) from None
raise
if filters.dataset_id and owner != filters.dataset_id:
raise HTTPException(502, "平台返回其他数据集字段")
for remote_key, key in (
("instrumentType", "instrument_type"),
("instrument_type", "instrument_type"),
("region", "region"),
("universe", "universe"),
("delay", "delay"),
):
if remote_key in row and row[remote_key] != getattr(filters, key):
raise HTTPException(502, "平台字段范围与查询不一致")
items.append(
encode_snapshot(
dict(
**data,
field_id=data["id"],
dataset_id=owner,
dataset_name=label(dataset) or owner,
source="worldquant",
collection_version=None,
scope=filters.model_dump(include=set(Scope.model_fields)),
fetched_at=now(),
synced_at=None,
)
)
)
count = raw.get("count")
known_total = type(count) is int and count >= 0
more = (
bool(raw["next"])
if "next" in raw
else (filters.offset + len(items) < count if known_total else len(items) == filters.limit)
)
result = page(
items,
count if known_total else filters.offset + len(items) + int(more),
filters.limit,
filters.offset,
)
result.update(has_more=more, total_known=known_total)
return result
async def resolve_fields(self, scope, refs):
"""Resolve trusted source records before any collection member is written."""
if any(ref.scope.key() != scope.key() for ref in refs):
raise HTTPException(422, "不能跨区域、Top、Delay 或品种添加字段")
result = {}
for ref in refs:
if ref.source == "local":
dataset = await Catalog(self.db).dataset(scope, ref.dataset_id, lock=True)
if not dataset.field_version or dataset.field_version != ref.collection_version:
raise HTTPException(409, "字段来源已更新,请重新查询后添加")
entry = await self.db.get(CatalogEntry, (dataset.field_version, ref.field_id))
if not entry:
raise HTTPException(422, "字段不属于指定数据集")
scope_row = await self.db.get(CatalogScope, scope.key())
owner = await self.db.get(CatalogEntry, (scope_row.catalog_version, ref.dataset_id))
field = self.local_field(entry, ref.dataset_id, owner.name if owner else None, scope)
else:
offset, seen, field = 0, set(), None
while True:
response = await self.online_fields(
FieldFilters(
**scope.model_dump(),
q=ref.field_id,
dataset_id=ref.dataset_id,
limit=100,
offset=offset,
)
)
field = next((item for item in response["items"] if item["id"] == ref.field_id), None)
if field or not response["has_more"]:
break
ids = {item["id"] for item in response["items"]}
if not ids - seen:
raise HTTPException(502, "平台字段分页未前进")
seen.update(ids)
offset += len(response["items"])
if not field:
raise HTTPException(422, "在线字段已不可用,请重新查询")
if field["id"] in result and result[field["id"]]["dataset_id"] != field["dataset_id"]:
raise HTTPException(422, "同名字段的数据集归属冲突")
result[field["id"]] = field
return list(result.values())
async def get(self, preparation_id, version=None, lock=False):
query = select(DataPreparation).where(DataPreparation.id == preparation_id)
row = await self.db.scalar(query.with_for_update() if lock else query)
if not row:
raise HTTPException(404, "数据准备集合不存在")
if version is not None and row.version != version:
raise HTTPException(409, "集合已修改,请重新读取或选择;当前草稿已保留")
return row
async def output(self, row):
count, datasets = (
await self.db.execute(
select(func.count(), func.count(func.distinct(PreparationField.dataset_id))).where(
PreparationField.preparation_id == row.id
)
)
).one()
return encode_snapshot(
dict(
id=row.id,
name=row.name,
note=row.note,
scope=row.scope,
version=row.version,
field_count=count,
dataset_count=datasets,
created_at=row.created_at,
updated_at=row.updated_at,
)
)
async def list(self, q="", scope_key=None, limit=25, offset=0):
query = select(DataPreparation)
if scope_key:
query = query.where(DataPreparation.scope_key == scope_key)
if q:
query = query.where(
or_(
DataPreparation.name.ilike(contains(q), escape="\\"),
DataPreparation.note.ilike(contains(q), escape="\\"),
)
)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(DataPreparation.updated_at.desc(), DataPreparation.id).limit(limit).offset(offset)
)
return page([await self.output(row) for row in rows], total, limit, offset)
async def members(self, preparation_id, q="", dataset_id=None, limit=25, offset=0):
await self.get(preparation_id)
query = select(PreparationField).where(PreparationField.preparation_id == preparation_id)
if dataset_id:
query = query.where(PreparationField.dataset_id == dataset_id)
if q:
query = query.where(
or_(
PreparationField.field_id.ilike(contains(q), escape="\\"),
PreparationField.content["name"].as_string().ilike(contains(q), escape="\\"),
PreparationField.content["description"].as_string().ilike(contains(q), escape="\\"),
)
)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(PreparationField.dataset_id, PreparationField.field_id).limit(limit).offset(offset)
)
return page([row.content for row in rows], total, limit, offset)
async def create(self, name, note, scope, fields):
row = DataPreparation(
id=str(uuid4()), name=name, note=note, scope=scope.model_dump(), scope_key=scope.key()
)
self.db.add(row)
await self.db.flush()
await self.add(row, fields)
return await self.output(row)
async def add(self, row, fields):
for field in fields:
existing = await self.db.get(PreparationField, (row.id, field["id"]))
if existing:
if existing.dataset_id != field["dataset_id"]:
raise HTTPException(422, "同名字段的数据集归属冲突")
continue
self.db.add(
PreparationField(
preparation_id=row.id, field_id=field["id"], dataset_id=field["dataset_id"], content=field
)
)
await self.db.flush()
async def copy_dataset(self, body):
dataset = await Catalog(self.db).dataset(body.scope, body.dataset_id, lock=True)
if not dataset.field_version or dataset.field_version != body.collection_version:
raise HTTPException(409, "数据集尚未完整同步或版本已更新")
batch = await self.db.get(CatalogBatch, dataset.field_version)
if not batch.complete:
raise HTTPException(409, "数据集尚未完整同步")
source = await Catalog(self.db).detail(body.scope, body.dataset_id)
rows = await self.db.scalars(
select(CatalogEntry)
.where(CatalogEntry.batch_id == dataset.field_version)
.order_by(CatalogEntry.id)
)
fields = [self.local_field(row, body.dataset_id, source["name"], body.scope) for row in rows]
return await self.create(
f"{source['name'] or body.dataset_id} · {now():%Y%m%d-%H%M%S-%f}", "", body.scope, fields
)
async def remove(self, refs):
rows = [await self.get(ref.id, ref.version, lock=True) for ref in sorted(refs, key=lambda r: r.id)]
for row in rows:
await self.db.execute(delete(PreparationField).where(PreparationField.preparation_id == row.id))
await self.db.delete(row)
return {"deleted": len(rows)}
async def freeze(self, refs):
"""Lock collection versions and capture source-independent research snapshots atomically."""
snapshots = []
for ref in sorted(refs, key=lambda r: r.id):
row = await self.get(ref.id, ref.version, lock=True)
existing = await self.db.scalar(
select(ResearchInputSnapshot).where(
ResearchInputSnapshot.preparation_id == row.id,
ResearchInputSnapshot.preparation_version == row.version,
)
)
if existing:
snapshots.append(await self.snapshot(existing.id))
continue
fields = [
r.content
for r in await self.db.scalars(
select(PreparationField)
.where(PreparationField.preparation_id == row.id)
.order_by(PreparationField.field_id)
)
]
if not fields:
raise HTTPException(422, "空集合不能用于研究")
fixed = ResearchInputSnapshot(
id=str(uuid4()),
preparation_id=row.id,
preparation_version=row.version,
content=dict(
name=row.name,
scope=row.scope,
fields=fields,
field_ids=[f["id"] for f in fields],
field_types={f["id"]: f["field_type"] for f in fields},
dataset_ids=sorted({f["dataset_id"] for f in fields}),
),
)
self.db.add(fixed)
await self.db.flush()
snapshots.append(await self.snapshot(fixed.id))
return snapshots
async def snapshot(self, snapshot_id):
row = await self.db.get(ResearchInputSnapshot, snapshot_id)
if not row:
raise HTTPException(404, "研究输入快照不存在")
return encode_snapshot(
dict(
**row.content,
id=row.id,
preparation_id=row.preparation_id,
preparation_version=row.preparation_version,
created_at=row.created_at,
)
)
async def bind(self, body):
refs = getattr(body, "preparation_refs", [])
if refs:
fixed = await self.freeze(refs)
body.input_ids = list(dict.fromkeys([*body.input_ids, *[r["id"] for r in fixed]]))
body.preparation_refs = []
if not body.input_ids:
raise HTTPException(422, "请选择非空的数据准备集合")
limit = next(
m.max_length for m in type(body).model_fields["input_ids"].metadata if hasattr(m, "max_length")
)
if len(body.input_ids) > limit:
raise HTTPException(422, f"最多可选择 {limit} 个研究输入")
return body
+16 -2
View File
@@ -4,6 +4,8 @@ from pydantic import Field
from ..ai.alpha_tools import AlphaArgs from ..ai.alpha_tools import AlphaArgs
from ..ai.capabilities import Capability from ..ai.capabilities import Capability
from ..preparations.service import Preparations
from ..schemas import Contract
from .contracts import ChatboxResearchInput, InputPageArgs, ResearchInputSelection, ResearchPreviewInput 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())) 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 = ( 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( Capability(
name="prepare_research_input", name="prepare_research_input",
schema=ResearchInputSelection, schema=ResearchInputSelection,
description="把明确选出的 1–100 个字段固定为研究输入快照;必须提供已读取的集合版本。只保存本地输入,不启动同步或回测。已有输入引用时直接读取,不重建。", description="将数据准备集合的明确 ID 和 version 固定为研究快照,保留完整字段与数据集归属。已有快照直接读取。",
label="固定研究输入", label="固定研究输入",
renderer="catalog", renderer="catalog",
effect="prepare", effect="prepare",
+5 -1
View File
@@ -57,7 +57,11 @@ class Assets:
"view": ViewSpec, "view": ViewSpec,
"workflow": WorkflowSpec, "workflow": WorkflowSpec,
}[body.kind] }[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": if body.kind == "workflow":
from .workflows import validate_graph from .workflows import validate_graph
+3 -6
View File
@@ -5,16 +5,13 @@ from typing import Literal
from pydantic import Field, model_validator from pydantic import Field, model_validator
from ..backtests.contracts import SimulationSettings, Source from ..backtests.contracts import SimulationSettings, Source
from ..catalog.contracts import Scope from ..preparations.contracts import PreparationReference
from ..schemas import Contract from ..schemas import Contract
from .expressions import PLACEHOLDER from .expressions import PLACEHOLDER
class ResearchInputSelection(Contract): class ResearchInputSelection(Contract):
scope: Scope items: list[PreparationReference] = Field(min_length=1, max_length=1)
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)
class InputPageArgs(Contract): class InputPageArgs(Contract):
@@ -50,7 +47,7 @@ class ChatboxResearchInput(Contract):
name: str = Field(min_length=1, max_length=200) name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=2000) 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) candidates: list[ResearchCandidate] = Field(min_length=1, max_length=100)
@model_validator(mode="after") @model_validator(mode="after")
+9 -10
View File
@@ -11,7 +11,8 @@ from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Simulatio
from ..backtests.service import Backtests, uid from ..backtests.service import Backtests, uid
from ..catalog.research_metadata import ResearchMetadata from ..catalog.research_metadata import ResearchMetadata
from ..catalog.service import Catalog 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 .assets import Assets
from .expressions import GROUPS, analyze, expand from .expressions import GROUPS, analyze, expand
from .serialization import encode_snapshot as jsonable_encoder from .serialization import encode_snapshot as jsonable_encoder
@@ -86,7 +87,7 @@ class Experiments:
"candidates": experiment["candidates"], "candidates": experiment["candidates"],
"hypothesis": experiment["hypothesis"], "hypothesis": experiment["hypothesis"],
"input_references": [ "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"] for entry in experiment["inputs"]
], ],
"template_reference": { "template_reference": {
@@ -133,6 +134,7 @@ class Experiments:
return validation return validation
async def create(self, body, kind="template", extra_evidence=None, *, parent_snapshots=None): 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 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 template = TemplateSpec.model_validate(asset["content"]) if asset else body.template
scope = scope_of(body.settings) scope = scope_of(body.settings)
@@ -318,7 +320,8 @@ class Experiments:
kind=source_kind or experiment["kind"], kind=source_kind or experiment["kind"],
reference=reference or experiment_id, reference=reference or experiment_id,
research_id=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], hypothesis=experiment["hypothesis"][:2000],
), ),
candidates=[ candidates=[
@@ -342,6 +345,7 @@ class Experiments:
original = parents[0] original = parents[0]
base = seed_settings(original["settings"]) base = seed_settings(original["settings"])
expression = original["expression"] expression = original["expression"]
await Preparations(self.db).bind(body)
snapshots, _ = await self.inputs(body.input_ids) snapshots, _ = await self.inputs(body.input_ids)
groups = defaultdict(list) groups = defaultdict(list)
for snapshot in snapshots: for snapshot in snapshots:
@@ -405,6 +409,7 @@ class Experiments:
) )
async def generation_context(self, body): async def generation_context(self, body):
await Preparations(self.db).bind(body)
snapshots, fields = await self.inputs(body.input_ids) snapshots, fields = await self.inputs(body.input_ids)
parents = await self.parents(body.parent_alpha_ids, body.parent_experiment_ids) parents = await self.parents(body.parent_alpha_ids, body.parent_experiment_ids)
metadata = await ResearchMetadata(self.db).operators(limit=100) metadata = await ResearchMetadata(self.db).operators(limit=100)
@@ -414,7 +419,7 @@ class Experiments:
"hypothesis": body.hypothesis, "hypothesis": body.hypothesis,
"method": body.method, "method": body.method,
"inputs": [ "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 for item in snapshots
], ],
"fields": dict(list(fields.items())[:300]), "fields": dict(list(fields.items())[:300]),
@@ -424,9 +429,3 @@ class Experiments:
], ],
"parents": [{**p, "candidates": p.get("candidates", [])[:10]} for p in parents], "parents": [{**p, "candidates": p.get("candidates", [])[:10]} for p in parents],
} }
async def available_inputs(self, limit=100):
rows = await self.db.scalars(
select(TemplateInput).order_by(TemplateInput.created_at.desc()).limit(limit)
)
return {"items": [await self.catalog.input(row.id) for row in rows]}
+2 -8
View File
@@ -28,12 +28,6 @@ from .workspace_contracts import (
router = APIRouter(prefix="/api/v1/research", tags=["research"], dependencies=[Depends(require_auth)]) router = 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") @router.get("/assets")
async def assets( async def assets(
request: Request, request: Request,
@@ -96,10 +90,10 @@ async def import_commit(body: ImportCommit, request: Request):
@router.post("/generate", status_code=201) @router.post("/generate", status_code=201)
async def generate(body: Generation, request: Request): 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) context = await Experiments(db).generation_context(body)
result, evidence = await request_model(request.app.state.ai, context, OUTPUTS[body.method]) 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, "模型不能改变已固定的输入范围") raise HTTPException(422, "模型不能改变已固定的输入范围")
async with request.app.state.sessions.begin() as db: async with request.app.state.sessions.begin() as db:
asset = await Assets(db).save( asset = await Assets(db).save(
+2 -2
View File
@@ -476,9 +476,9 @@ class ResearchRuntime:
step = await db.get(ResearchStepRun, step_id) step = await db.get(ResearchStepRun, step_id)
if not step or step.status != "running": if not step or step.status != "running":
return 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"]] [i["id"] for i in step.output["context"]["inputs"]]
): )):
raise HTTPException(422, "模型不能改变已固定的输入范围") raise HTTPException(422, "模型不能改变已固定的输入范围")
# A paused/stopped run may collect this already-issued model output, but cannot advance. # A paused/stopped run may collect this already-issued model output, but cannot advance.
asset = await Assets(db).save( asset = await Assets(db).save(
+10 -50
View File
@@ -5,12 +5,9 @@ not FASTEXPR operator semantics or the account's current platform permissions.
""" """
from fastapi import HTTPException from fastapi import HTTPException
from sqlalchemy import select
from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Source from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Source
from ..catalog.contracts import EntryOutput, InputPreparation
from ..catalog.service import Catalog from ..catalog.service import Catalog
from ..models import CatalogEntry
from .expressions import analyze, expand from .expressions import analyze, expand
@@ -21,55 +18,18 @@ class ResearchBuilder:
self.backtests = backtests self.backtests = backtests
async def select_input(self, body): async def select_input(self, body):
"""Fix explicit fields in one published version; reject missing or stale members.""" from ..preparations.service import Preparations
collection = await self.catalog.collection(body.scope, body.dataset_id) saved = (await Preparations(self.db).freeze(body.items))[0]
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],
)
)
return await self.input_page(saved["id"]) return await self.input_page(saved["id"])
async def input_page(self, input_id, limit=25, offset=0, q="", field_type=None): 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) saved = await self.catalog.input(input_id)
ids = [ fields = [f for f in saved["fields"] if (not q or q.lower() in
field " ".join(str(f.get(k) or "") for k in ("id", "name", "description", "dataset_id")).lower())
for field in saved["field_ids"] and (not field_type or f["field_type"] == field_type)]
if q.lower() in field.lower() return {**{k: v for k, v in saved.items() if k not in ("fields", "field_ids", "field_types")},
and (field_type is None or saved["field_types"].get(field) == field_type) "field_count": len(saved["fields"]), "items": fields[offset:offset + limit],
] "total": len(fields), "limit": limit, "offset": offset, "has_more": offset + limit < len(fields)}
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),
}
async def prepare(self, body): async def prepare(self, body):
"""Bind templates against an immutable input, then reuse the fixed-preview interface. """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. Raises HTTPException(422) for wrong scope, membership or declared type.
No expression execution or implicit cleaning/aggregation takes place here. 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"] scope = saved["scope"]
candidates = [] candidates = []
for item in body.candidates: for item in body.candidates:
@@ -114,7 +74,7 @@ class ResearchBuilder:
source = Source.model_validate( source = Source.model_validate(
{ {
**body.source.model_dump(), **body.source.model_dump(),
"template_input_id": saved["id"], "input_snapshot_id": saved["id"],
"hypothesis": body.hypothesis, "hypothesis": body.hypothesis,
} }
) )
+2
View File
@@ -181,6 +181,8 @@ class Workflows:
for node in graph.nodes: for node in graph.nodes:
if node.type == "iterate" and node.config["max_rounds"] > body.budget.max_rounds: if node.type == "iterate" and node.config["max_rounds"] > body.budget.max_rounds:
raise HTTPException(422, "流程迭代上限超过本次授权轮数") raise HTTPException(422, "流程迭代上限超过本次授权轮数")
from ..preparations.service import Preparations
await Preparations(self.db).bind(body)
experiments = Experiments(self.db) experiments = Experiments(self.db)
settings_variant = any( settings_variant = any(
n.type == "variant" and n.config.get("method") == "settings" for n in graph.nodes n.type == "variant" and n.config.get("method") == "settings" for n in graph.nodes
+11 -5
View File
@@ -7,6 +7,7 @@ from pydantic import Field, field_validator, model_validator
from ..backtests.contracts import SimulationSettings from ..backtests.contracts import SimulationSettings
from ..catalog.contracts import Scope from ..catalog.contracts import Scope
from ..preparations.contracts import PreparationReference
from ..schemas import Contract from ..schemas import Contract
from .expressions import IDENTIFIER, PLACEHOLDER, normalize_template from .expressions import IDENTIFIER, PLACEHOLDER, normalize_template
@@ -70,7 +71,8 @@ class FeatureStep(Contract):
class FeatureSpec(Contract): class FeatureSpec(Contract):
name: str = Field(min_length=1, max_length=200) name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=10000) 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) steps: list[FeatureStep] = Field(default_factory=list, max_length=30)
template: TemplateSpec | None = None template: TemplateSpec | None = None
@@ -102,7 +104,8 @@ class Expansion(Contract):
asset_id: str | None = Field(default=None, max_length=36) asset_id: str | None = Field(default=None, max_length=36)
version: int | None = Field(default=None, ge=1) version: int | None = Field(default=None, ge=1)
template: TemplateSpec | None = None 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) hypothesis: str = Field(min_length=1, max_length=10000)
settings: SimulationSettings settings: SimulationSettings
mode: Literal["all", "random"] = "all" mode: Literal["all", "random"] = "all"
@@ -123,7 +126,8 @@ class Expansion(Contract):
class Generation(Contract): class Generation(Contract):
name: str = Field(min_length=1, max_length=200) name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=10000) 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_alpha_ids: list[str] = Field(default_factory=list, max_length=20)
parent_experiment_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" method: Literal["template", "structure", "feature"] = "template"
@@ -131,7 +135,8 @@ class Generation(Contract):
class SettingVariants(Contract): class SettingVariants(Contract):
alpha_id: str = Field(min_length=1, max_length=100) 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) hypothesis: str = Field(default="保持表达式,比较字段共同支持的市场与设置", max_length=10000)
@@ -229,7 +234,8 @@ class FlowStart(Contract):
name: str = Field(min_length=1, max_length=200) name: str = Field(min_length=1, max_length=200)
workflow_id: str | None = None workflow_id: str | None = None
workflow_version: int | None = Field(default=None, ge=1) 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) hypothesis: str = Field(min_length=1, max_length=10000)
settings: SimulationSettings settings: SimulationSettings
budget: Budget budget: Budget
+17
View File
@@ -7,6 +7,7 @@ from pydantic import Field, model_validator
from ..backtests.contracts import Candidate, SimulationSettings from ..backtests.contracts import Candidate, SimulationSettings
from ..catalog.contracts import CatalogFilters, Scope from ..catalog.contracts import CatalogFilters, Scope
from ..preparations.contracts import PreparationReference
from ..research.workspace_contracts import TemplateSpec from ..research.workspace_contracts import TemplateSpec
from ..schemas import Contract from ..schemas import Contract
@@ -54,8 +55,24 @@ class Provenance(Contract):
parent_run_id: RunId | None = None 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): class Submit(Contract):
name: str = Field(min_length=1, max_length=200) 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) candidates: list[DirectCandidate] = Field(min_length=1, max_length=100)
idempotency_key: Identifier idempotency_key: Identifier
duplicate_policy: Literal["reject", "rerun"] = "reject" duplicate_policy: Literal["reject", "rerun"] = "reject"
+13 -1
View File
@@ -118,6 +118,17 @@ class ResearchAccess:
return {"job_id": job.id, "status": job.status, "action": job.kind, return {"job_id": job.id, "status": job.status, "action": job.kind,
"read_with": "get_worldquant_connection", "web_url": f"{self.public_origin}/"} "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): async def catalog(self, args):
data = await Catalog(self.db).search(args.filters, args.dataset_id) data = await Catalog(self.db).search(args.filters, args.dataset_id)
return {**data, "scope": args.filters.model_dump(include={"region", "universe", "delay", "instrument_type"}), 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. # preserve_source prevents the Chatbox-specific generating context rewriting MCP provenance.
backtests = Backtests(self.db, provenance) backtests = Backtests(self.db, provenance)
preview = await backtests.preview(PreviewInput(inline=DraftInput( 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"], result = await backtests.start(StartInput(preview_id=preview["preview_id"],
idempotency_key="mcp-" + str(uuid4()))) idempotency_key="mcp-" + str(uuid4())))
result = {**result, "input_digest": digest, "batch_count": preview["batch_count"], result = {**result, "input_digest": digest, "batch_count": preview["batch_count"],
@@ -0,0 +1,75 @@
"""Replace saved dataset inputs with editable preparations and independent snapshots."""
import sqlalchemy as sa
from alembic import op
revision = "0015"
down_revision = "0014"
branch_labels = None
depends_on = None
def upgrade():
op.drop_table("template_inputs")
op.add_column("catalog_batches", sa.Column("job_id", sa.String(36), nullable=True))
op.add_column("catalog_batches", sa.Column("offset", sa.Integer(), nullable=False, server_default="0"))
op.execute("UPDATE catalog_batches SET job_id = id")
# Preserve unfinished catalog pagination; only legacy research inputs are discarded.
jobs = sa.table("sync_jobs", sa.column("id"), sa.column("checkpoint", sa.JSON()))
batches = sa.table("catalog_batches", sa.column("id"), sa.column("offset", sa.Integer()))
for job_id, checkpoint in op.get_bind().execute(sa.select(jobs.c.id, jobs.c.checkpoint)):
offset = (checkpoint or {}).get("offset", 0)
if type(offset) is int and offset >= 0:
op.get_bind().execute(batches.update().where(batches.c.id == job_id).values(offset=offset))
names = {"fk": "fk_%(table_name)s_%(column_0_name)s_%(referred_table_name)s"}
fk = next(
f
for f in sa.inspect(op.get_bind()).get_foreign_keys("catalog_batches")
if f["constrained_columns"] == ["id"]
)
with op.batch_alter_table("catalog_batches", naming_convention=names) as batch:
batch.drop_constraint(fk["name"] or "fk_catalog_batches_id_sync_jobs", type_="foreignkey")
batch.alter_column("job_id", existing_type=sa.String(36), nullable=False)
batch.create_foreign_key("fk_catalog_batches_job", "sync_jobs", ["job_id"], ["id"])
batch.create_index("ix_catalog_batches_job_id", ["job_id"])
op.create_table(
"data_preparations",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("name", sa.String(200), nullable=False),
sa.Column("note", sa.Text(), nullable=False),
sa.Column("scope_key", sa.String(200), nullable=False),
sa.Column("scope", sa.JSON(), nullable=False),
sa.Column("version", sa.Integer(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
)
op.create_index("ix_data_preparations_scope_key", "data_preparations", ["scope_key"])
op.create_table(
"preparation_fields",
sa.Column(
"preparation_id",
sa.String(36),
sa.ForeignKey("data_preparations.id", ondelete="CASCADE"),
primary_key=True,
),
sa.Column("field_id", sa.String(200), primary_key=True),
sa.Column("dataset_id", sa.String(200), nullable=False),
sa.Column("content", sa.JSON(), nullable=False),
)
op.create_index("ix_preparation_fields_dataset_id", "preparation_fields", ["dataset_id"])
op.create_table(
"research_input_snapshots",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("preparation_id", sa.String(36), nullable=False),
sa.Column("preparation_version", sa.Integer(), nullable=False),
sa.Column("content", sa.JSON(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.UniqueConstraint("preparation_id", "preparation_version"),
)
op.create_index(
"ix_research_input_snapshots_preparation_id", "research_input_snapshots", ["preparation_id"]
)
def downgrade():
raise RuntimeError("旧输入模型已移除;回退请恢复升级前数据库备份")
+3 -5
View File
@@ -43,8 +43,7 @@ if __name__ == "__main__":
"INSERT INTO research (alpha_id, note, tags, favorite, state, updated_at, version) VALUES ('MIGRATION_TEST', 'preserve research', '[]', false, 'inbox', now(), 7);" "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.upgrade(config, "0014")
command.check(config)
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [ assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
("preserve research", 7) ("preserve research", 7)
] ]
@@ -56,7 +55,7 @@ if __name__ == "__main__":
("preserve research", 7) ("preserve research", 7)
] ]
print( 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(): async def flow():
@@ -73,7 +72,6 @@ if __name__ == "__main__":
return httpx.Response(201, json={"token": {"expiry": 14400}}) return httpx.Response(201, json={"token": {"expiry": 14400}})
if request.url.path == "/users/self": if request.url.path == "/users/self":
return httpx.Response(200, json={"id": "PG_TEST_USER"}) return httpx.Response(200, json={"id": "PG_TEST_USER"})
assert request.method == "GET"
return catalog_response(request) or httpx.Response(404) return catalog_response(request) or httpx.Response(404)
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver") 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] assert sorted(r.status_code for r in responses) == [200, 409]
await sync(catalog, "TEST_FIN") await sync(catalog, "TEST_FIN")
assert (await prepare(client, version)).status_code == 409 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 assert persisted == draft
print( print(
"PostgreSQL: real API/runner multi-page dedupe, immutable draft, refresh conflict and concurrent note CAS passed" "PostgreSQL: real API/runner multi-page dedupe, immutable draft, refresh conflict and concurrent note CAS passed"
+139
View File
@@ -0,0 +1,139 @@
"""Isolated PostgreSQL migration and concurrency acceptance, using synthetic upstream only."""
import asyncio
import os
from alembic import command
from alembic.config import Config
from cryptography.fernet import Fernet
from sqlalchemy import text
from sqlalchemy.ext.asyncio import create_async_engine
URL = "postgresql+asyncpg://postgres:preparations-test-only@127.0.0.1:18437/preparations_test"
os.environ.update(
DATABASE_URL=URL,
ADMIN_PASSWORD="migration-test-only",
ENCRYPTION_KEY=Fernet.generate_key().decode(),
WQ_EMAIL="",
WQ_PASSWORD="",
)
async def sql(query):
engine = create_async_engine(URL)
try:
async with engine.begin() as db:
result = await db.execute(text(query))
return result.fetchall() if result.returns_rows else None
finally:
await engine.dispose()
async def acceptance():
import httpx
from app.catalog.contracts import CatalogJobInput, Scope
from app.catalog.service import Catalog
from app.config import Settings
from app.main import create_app
from app.worldquant import WqClient
from tests.catalog_fake import catalog_response
from tests.test_catalog import SCOPE, prepare, search, sync
def upstream(request):
if request.url.path == "/authentication":
return httpx.Response(201, json={"token": {"expiry": 14400}})
if request.url.path == "/users/self":
return httpx.Response(200, json={"id": "PG_TEST_USER"})
from tests.research_metadata_fake import response
metadata = response(request)
if metadata is not None:
return metadata
assert request.method == "GET"
return catalog_response(request) or httpx.Response(404)
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
app = create_app(settings, WqClient(settings, transport=httpx.MockTransport(upstream)))
async with app.router.lifespan_context(app):
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
headers={"X-WQ-Request": "1"},
) as client:
assert (
await client.post(
"/api/v1/auth/login", json={"username": "admin", "password": "migration-test-only"}
)
).status_code == 200
await client.put(
"/api/v1/account/credentials",
json={"email": "synthetic@example.com", "password": "synthetic-only"},
)
job = (await client.post("/api/v1/account/connect")).json()
await app.state.runner.execute(job["id"])
fixture = (client, app.state.runner, {})
await sync(fixture)
await sync(fixture, "TEST_FIN")
snapshot = (
await prepare(
client, (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"]
)
).json()
assert len(snapshot["field_ids"]) == 123
fields = await client.get("/api/v1/catalog/fields", params={**SCOPE, "category": "基本面", "limit": 2, "offset": 2})
assert fields.status_code == 200, fields.text
assert fields.json()["total"] == 123 and len(fields.json()["items"]) == 2
assert fields.json()["items"][0]["category"] == "基本面"
async def enqueue():
async with app.state.sessions.begin() as db:
return (await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)).id
jobs = await asyncio.gather(*(enqueue() for _ in range(5)))
assert len(set(jobs)) == 1, jobs
collection = (await client.get("/api/v1/data-preparations")).json()["items"][0]
ref = {"id": collection["id"], "version": collection["version"]}
async def freeze():
response = await client.post("/api/v1/data-preparations/freeze", json={"items": [ref]})
assert response.status_code == 201, response.text
return response.json()["items"][0]["id"]
assert len(set(await asyncio.gather(*(freeze() for _ in range(5))))) == 1
await client.delete(f"/api/v1/data-preparations/{ref['id']}?version={ref['version']}")
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
if __name__ == "__main__":
assert not asyncio.run(sql("SELECT tablename FROM pg_tables WHERE schemaname='public'")), (
"Requires an empty isolated database"
)
config = Config("alembic.ini")
command.upgrade(config, "0014")
# Existing catalog checkpoints survive the structural change; no research data migration.
asyncio.run(
sql(
"INSERT INTO sync_jobs(id,kind,status,payload,checkpoint,total,processed,failed,cancel_requested,created_at,updated_at) VALUES ('checkpoint-test','catalog_sync','failed','{}','{\"offset\": 100}',0,100,0,false,now(),now())"
)
)
asyncio.run(sql("INSERT INTO catalog_scopes(key,scope) VALUES ('checkpoint-scope','{}')"))
asyncio.run(
sql(
"INSERT INTO catalog_batches(id,scope_key,dataset_id,complete,count) VALUES ('checkpoint-test','checkpoint-scope',NULL,false,100)"
)
)
command.upgrade(config, "head")
command.check(config)
assert asyncio.run(
sql("SELECT job_id, catalog_batches.\"offset\" FROM catalog_batches WHERE id='checkpoint-test'")
) == [("checkpoint-test", 100)]
assert asyncio.run(sql("SELECT to_regclass('template_inputs')")) == [(None,)]
asyncio.run(sql("DELETE FROM catalog_batches WHERE id='checkpoint-test'"))
asyncio.run(sql("DELETE FROM catalog_scopes WHERE key='checkpoint-scope'"))
asyncio.run(sql("DELETE FROM sync_jobs WHERE id='checkpoint-test'"))
asyncio.run(acceptance())
print(
"PostgreSQL 17: 0014 → 0015 metadata, retained catalog checkpoint, concurrent job deduplication/freeze and independent snapshot passed"
)
+8 -24
View File
@@ -22,7 +22,7 @@ def research_step(text, returns, history):
return "get_backtest_results", {"run_id": run_id} return "get_backtest_results", {"run_id": run_id}
data = content(returns[-1]) data = content(returns[-1])
return f"已读取本次研究的 {data['total']} 条真实保存结果;缺失 Sharpe 仍为未知。" 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 "请先保存字段选择,再点击用此输入研究。" return "请先保存字段选择,再点击用此输入研究。"
if not returns: if not returns:
return "get_backtest_capabilities", {} return "get_backtest_capabilities", {}
@@ -30,39 +30,23 @@ def research_step(text, returns, history):
data = content(last) data = content(last)
if "error" in data: if "error" in data:
return f"研究尚未完成:{data['error']}" 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 last.tool_name == "get_backtest_capabilities":
if context.get("template_input_id"): if context.get("input_snapshot_id"):
return "get_research_input", { return "get_research_input", {
"input_id": context["template_input_id"], "input_id": context["input_snapshot_id"],
"field_type": "MATRIX", "field_type": "MATRIX",
"limit": 1, "limit": 1,
} }
return "search_catalog", {"filters": {**scope, "q": "TEST_FIN", "limit": 1}} return "search_data_preparations", {"limit": 1}
if last.tool_name == "search_catalog": if last.tool_name == "search_data_preparations":
if data["dataset_id"] is None: return "prepare_research_input", {"items": [{"id": data["items"][0]["id"], "version": data["items"][0]["version"]}]}
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"]],
}
if last.tool_name in ("get_research_input", "prepare_research_input"): 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"] saved_scope = data["scope"]
return "prepare_research_backtest", { return "prepare_research_backtest", {
"name": "Chatbox 数据集研究", "name": "Chatbox 数据集研究",
"hypothesis": "验证所选合成字段的横截面排序信号", "hypothesis": "验证所选合成字段的横截面排序信号",
"template_input_id": data["id"], "input_snapshot_id": data["id"],
"candidates": [ "candidates": [
{ {
"client_item_id": "research-1", "client_item_id": "research-1",
+2 -2
View File
@@ -30,7 +30,7 @@ async def acceptance():
from app.config import Settings from app.config import Settings
from app.main import create_app from app.main import create_app
from app.models import TemplateInput from app.models import ResearchInputSnapshot
from app.research.workspace_contracts import TemplateSpec from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure from tests.test_ai import configure
from tests.test_backtests import setup from tests.test_backtests import setup
@@ -66,7 +66,7 @@ async def acceptance():
} }
async with app.state.sessions() as db: async with app.state.sessions() as db:
fixed = await db.scalar(select(TemplateInput)) fixed = await db.scalar(select(ResearchInputSnapshot))
body = { body = {
"request_id": "finite-run", "request_id": "finite-run",
"name": "PG 有限研究", "name": "PG 有限研究",
+3 -2
View File
@@ -170,14 +170,15 @@ async def main(args):
fields = await api("GET", f"/catalog/datasets/pv1/fields?{params}&limit=100") fields = await api("GET", f"/catalog/datasets/pv1/fields?{params}&limit=100")
fixed_input = await api( fixed_input = await api(
"POST", "POST",
"/catalog/inputs", "/data-preparations/from-dataset",
{ {
"scope": scope, "scope": scope,
"dataset_id": "pv1", "dataset_id": "pv1",
"collection_version": fields["collection_version"], "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( availability = await api(
"POST", "/catalog/field-availability/refresh", {"field_id": "close", "scope": scope} "POST", "/catalog/field-availability/refresh", {"field_id": "close", "scope": scope}
) )
+2 -2
View File
@@ -28,7 +28,7 @@ async def acceptance():
from app.config import Settings from app.config import Settings
from app.main import create_app 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 ( from tests.test_research_outcomes import (
test_feature_conversion_keeps_original_version_through_experiment, test_feature_conversion_keeps_original_version_through_experiment,
test_lineage_retains_multiple_parents_and_descendants, test_lineage_retains_multiple_parents_and_descendants,
@@ -46,7 +46,7 @@ async def acceptance():
) )
assert response.status_code == 200 assert response.status_code == 200
async with app.state.sessions() as db: 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 experiment in await db.scalars(select(ResearchExperiment)):
for parent in experiment.parents: for parent in experiment.parents:
assert await db.get(ResearchParent, (experiment.id, parent["kind"], parent["id"])) assert await db.get(ResearchParent, (experiment.id, parent["kind"], parent["id"]))
+2 -2
View File
@@ -30,7 +30,7 @@ async def acceptance():
from app.config import Settings from app.config import Settings
from app.main import create_app from app.main import create_app
from app.models import TemplateInput from app.models import ResearchInputSnapshot
from app.research.workspace_contracts import TemplateSpec from app.research.workspace_contracts import TemplateSpec
from tests.test_ai import configure from tests.test_ai import configure
from tests.test_backtests import setup from tests.test_backtests import setup
@@ -67,7 +67,7 @@ async def acceptance():
} }
async with app.state.sessions() as db: async with app.state.sessions() as db:
fixed = await db.scalar(select(TemplateInput)) fixed = await db.scalar(select(ResearchInputSnapshot))
body = { body = {
"request_id": "finite-run", "request_id": "finite-run",
"name": "PG 有限研究", "name": "PG 有限研究",
+6 -9
View File
@@ -10,11 +10,10 @@ from sqlalchemy import func, select
from app.ai.capabilities import ToolContext, assemble from app.ai.capabilities import ToolContext, assemble
from app.ai.tools import CAPABILITIES from app.ai.tools import CAPABILITIES
from app.business import Business 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 app.research.service import ResearchBuilder
from tests.test_ai import configure, single_tool_factory, start from tests.test_ai import configure, single_tool_factory, start
from tests.test_api import seed from tests.test_api import seed
from tests.test_catalog import SCOPE
from tests.test_catalog import catalog as catalog_fixture from tests.test_catalog import catalog as catalog_fixture
from tests.test_research_integration import fixed_input as fixed_input_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. # select_input has already persisted the new input before requesting its result page.
raise HTTPException(422, "准备输入后的校验失败") 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) monkeypatch.setattr(ResearchBuilder, "input_page", unavailable_page)
app.state.ai.model_factory = single_tool_factory( app.state.ai.model_factory = single_tool_factory(
"prepare_research_input", "prepare_research_input",
{ {"items": [{"id": fixed_input["preparation_id"], "version": updated.json()["version"]}]},
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
) )
_, run, _ = await start(app, logged_in, "保存研究输入") _, run, _ = await start(app, logged_in, "保存研究输入")
call = run["tools"][0] 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["presentation"]["effect"] == "prepare"
assert call["result"]["error"] == "准备输入后的校验失败" assert call["result"]["error"] == "准备输入后的校验失败"
async with app.state.sessions() as db: 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" assert (await db.get(AIToolCall, call["id"])).status == "failed"
+3 -2
View File
@@ -212,8 +212,7 @@ def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, m
}, },
) )
for _ in range(2): for _ in range(2):
command.upgrade(config, "head") command.upgrade(config, "0014")
command.check(config)
alphas = sa.Table("alphas", sa.MetaData(), autoload_with=engine) alphas = sa.Table("alphas", sa.MetaData(), autoload_with=engine)
with engine.connect() as db: with engine.connect() as db:
rows = db.execute(sa.select(alphas).order_by(alphas.c.id)).mappings().all() 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() record = db.execute(sa.select(research)).mappings().one()
assert record["tags"] == ["PPAC"] and record["note"] == "keep" and record["version"] == 7 assert record["tags"] == ["PPAC"] and record["note"] == "keep" and record["version"] == 7
command.downgrade(config, "0010") command.downgrade(config, "0010")
command.upgrade(config, "head")
command.check(config)
engine.dispose() engine.dispose()
+1
View File
@@ -203,6 +203,7 @@ async def test_export_formula_injection_and_detail_variants(app, logged_in):
assert row["name"] == "'=DANGEROUS()" and row["note"] == "' @formula()" assert row["name"] == "'=DANGEROUS()" and row["note"] == "' @formula()"
assert (await logged_in.get(f"{PREFIX}/alphas/super1/pnl")).json() == { assert (await logged_in.get(f"{PREFIX}/alphas/super1/pnl")).json() == {
"cached": False, "cached": False,
"series": [],
"points": [], "points": [],
"fetched_at": None, "fetched_at": None,
} }
+21 -13
View File
@@ -86,16 +86,24 @@ async def search(client, suffix="/datasets", **params):
async def prepare(client, version, **changes): async def prepare(client, version, **changes):
return await client.post( response = await client.post("/api/v1/data-preparations/from-dataset", json={
BASE + "/inputs", "scope": changes.get("scope", SCOPE), "dataset_id": changes.get("dataset_id", "TEST_FIN"),
json={ "collection_version": version,
"scope": SCOPE, })
"dataset_id": "TEST_FIN", if response.status_code != 201:
"collection_version": version, return response
"selection": "all", collection = response.json()
**changes, 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): 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) response = await prepare(client, version)
assert response.status_code == 201, response.text assert response.status_code == 201, response.text
draft = response.json() 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"]: for suffix in ["/datasets/TEST_FIN", "/datasets/TEST_FIN/fields/TEST_FIN_122"]:
detail = await search(client, suffix) detail = await search(client, suffix)
assert detail["research"]["version"] == 1 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 newer["collection_version"] != version
assert (await prepare(client, version)).status_code == 409 assert (await prepare(client, version)).status_code == 409
assert len((await prepare(client, newer["collection_version"])).json()["field_ids"]) == 125 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"][ assert (await search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122"))["research"][
"note" "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, scope=other)
await sync(catalog, "TEST_FIN", scope=other) await sync(catalog, "TEST_FIN", scope=other)
assert (await prepare(client, version, scope=other)).status_code == 409 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): async def test_catalog_authentication_and_origin(app, client):
+26 -1
View File
@@ -162,7 +162,8 @@ async def test_official_sdk_client_and_error_contract(mcp_app):
async with ClientSession(streams[0], streams[1]) as client: async with ClientSession(streams[0], streams[1]) as client:
await client.initialize() await client.initialize()
listed = await client.list_tools() 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", {}) caps = await client.call_tool("get_research_capabilities", {})
assert caps.structured_content["max_candidates"] == 100 assert caps.structured_content["max_candidates"] == 100
result = await client.call_tool("submit_backtests", submission()) 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" assert missing.structured_content["error"]["code"] == "CREDENTIALS_NOT_CONFIGURED"
wrong = await mcp_app.state.mcp.invoke(principal, "get_worldquant_connection", {"job_id": "unrelated"}) wrong = await mcp_app.state.mcp.invoke(principal, "get_worldquant_connection", {"job_id": "unrelated"})
assert wrong.structured_content["error"]["code"] == "NOT_FOUND" assert wrong.structured_content["error"]["code"] == "NOT_FOUND"
async def test_mcp_preparations_freeze_at_submit_and_survive_deletion(mcp_app):
from app.catalog.contracts import Scope
from app.preparations.contracts import PreparationReference
from app.preparations.service import Preparations
principal, _ = await credentials(mcp_app)
async with mcp_app.state.sessions.begin() as db:
collection = await Preparations(db).create("MCP prepared", "", Scope(region="USA", universe="TOP3000", delay=1),
[{"id": "close", "field_id": "close", "name": "Close", "dataset_id": "pv1", "dataset_name": "Price",
"description": "Synthetic close", "field_type": "MATRIX", "source": "local", "fetched_at": now().isoformat()}])
refs = [{"id": collection["id"], "version": 1}]
found = await invoke(mcp_app, principal, "search_data_preparations", {"q": "MCP prepared"})
assert found["items"][0]["id"] == collection["id"]
detail = await invoke(mcp_app, principal, "get_data_preparation", {**refs[0], "limit": 1})
assert detail["fields"]["items"][0]["dataset_id"] == "pv1"
result = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
async with mcp_app.state.sessions.begin() as db:
run = await db.get(BacktestRun, result["backtest_run_id"])
snapshot_id = run.source["input_snapshot_ids"][0]
await Preparations(db).remove([PreparationReference(**refs[0])])
assert (await Preparations(db).snapshot(snapshot_id))["fields"][0]["description"] == "Synthetic close"
# Idempotent replay uses the already fixed run even after the collection is gone.
replay = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
assert replay["backtest_run_id"] == result["backtest_run_id"]
+352
View File
@@ -0,0 +1,352 @@
"""Preparation contracts across real HTTP, transactions and immutable research inputs."""
from argparse import Namespace
import pytest
from sqlalchemy import select
from app.catalog.contracts import CatalogJobInput, Scope
from app.catalog.service import Catalog
from app.models import CatalogBatch, Job, ResearchInputSnapshot
from tests.catalog_fake import field_records
from tests.test_catalog import SCOPE, sync
from tests.test_catalog import catalog as catalog_fixture
catalog = catalog_fixture
BASE = "/api/v1/data-preparations"
def reference(field):
return {k: field[k] for k in ("scope", "dataset_id", "field_id", "source", "collection_version")}
async def copied(catalog):
client = catalog[0]
await sync(catalog)
job = await sync(catalog, "TEST_FIN")
response = await client.post(
BASE + "/from-dataset",
json={"scope": SCOPE, "dataset_id": "TEST_FIN", "collection_version": job["id"]},
)
assert response.status_code == 201, response.text
return response.json()
async def test_multi_dataset_collection_atomic_changes_and_snapshot_independence(catalog):
client, runner, state = catalog
collection = await copied(catalog)
state["fields"] = field_records("TEST_NEWS", 3)
await sync(catalog, "TEST_NEWS")
query = (await client.get("/api/v1/catalog/fields", params={**SCOPE, "dataset_id": "TEST_NEWS"})).json()
assert query["total"] == 3
ref = reference(query["items"][0])
response = await client.patch(
f"{BASE}/{collection['id']}/fields", json={"version": 1, "fields": [ref, ref]}
)
assert response.status_code == 200, response.text
collection = response.json()
assert collection["field_count"] == 124 and collection["dataset_count"] == 2
bad = {**ref, "scope": {**SCOPE, "delay": 0}}
response = await client.patch(
f"{BASE}/{collection['id']}/fields",
json={"version": 2, "fields": [ref, bad], "remove_ids": ["TEST_FIN_001"]},
)
assert response.status_code == 422
assert (await client.get(f"{BASE}/{collection['id']}")).json()["version"] == 2
snapshot = (
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
).json()["items"][0]
assert snapshot["dataset_ids"] == ["TEST_FIN", "TEST_NEWS"]
assert all("description" in f and f["dataset_id"] for f in snapshot["fields"])
assert (
await client.patch(f"{BASE}/{collection['id']}", json={"version": 2, "name": "edited"})
).status_code == 200
assert (
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
).status_code == 409
assert (await client.delete(f"{BASE}/{collection['id']}?version=3")).status_code == 200
assert (await client.get(f"{BASE}/{collection['id']}")).status_code == 404
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
async with runner.sessions() as db:
assert await db.get(ResearchInputSnapshot, snapshot["id"])
@pytest.mark.parametrize(
"key,value", [("region", "CHN"), ("universe", "TOP1000"), ("delay", 0), ("instrument_type", "FUTURE")]
)
async def test_scope_dimensions_rejected_without_creating_collection(catalog, key, value):
client = catalog[0]
await copied(catalog)
field = (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["items"][0]
before = (await client.get(BASE)).json()["total"]
response = await client.post(
BASE, json={"name": "bad", "scope": {**SCOPE, key: value}, "fields": [reference(field)]}
)
assert response.status_code == 422
assert (await client.get(BASE)).json()["total"] == before
async def test_online_fields_without_sync_and_local_search_before_pagination(catalog):
client, _, _ = catalog
response = await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100})
assert response.status_code == 200, response.text
first = response.json()
assert len(first["items"]) == 100 and first["has_more"]
second = (
await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100, "offset": 100})
).json()
assert len(second["items"]) > 0
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
field = reference(second["items"][-1])
response = await client.post(BASE, json={"name": "online only", "scope": SCOPE, "fields": [field]})
assert response.status_code == 201, response.text
assert response.json()["field_count"] == 1
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
await copied(catalog)
query = (
await client.get(
"/api/v1/catalog/fields", params={**SCOPE, "q": "合成字段说明 11", "limit": 2, "offset": 2}
)
).json()
assert query["total"] == 11 and len(query["items"]) == 2
response = await client.get(
"/api/v1/catalog/fields", params={**SCOPE, "coverage_min": 0.9, "coverage_max": 0.5}
)
assert response.status_code == 422
async def test_empty_crud_copy_and_batch_delete_are_version_checked(catalog):
client = catalog[0]
first = (await client.post(BASE, json={"name": "empty", "scope": SCOPE})).json()
assert first["field_count"] == 0
assert (
await client.post(BASE + "/freeze", json={"items": [{"id": first["id"], "version": 1}]})
).status_code == 422
second = (await client.post(f"{BASE}/{first['id']}/copy", json={"version": 1})).json()
response = await client.post(
BASE + "/batch-delete",
json={"items": [{"id": first["id"], "version": 1}, {"id": second["id"], "version": 99}]},
)
assert response.status_code == 409
assert (await client.get(BASE)).json()["total"] == 2
assert (
await client.patch(
f"{BASE}/{first['id']}", json={"name": "x", "version": 1, "scope": {**SCOPE, "delay": 0}}
)
).status_code == 422
async def test_full_sync_partial_failure_and_resume_publish_per_dataset(catalog):
client, runner, state = catalog
await copied(catalog)
async with runner.sessions.begin() as db:
service = Catalog(db)
job = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
duplicate = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
assert duplicate.id == job.id
# The fixture deliberately returns TEST_FIN ownership for the other datasets.
await runner.execute(job.id)
failed = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert failed["status"] == "completed_with_errors" and failed["processed"] == 1
version = (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
"collection_version"
]
assert version != job.id
async with runner.sessions() as db:
assert (await db.get(CatalogBatch, version)).job_id == job.id
state["fields"] = field_records("TEST_NEWS", 3)
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
await runner.execute(job.id)
retried = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert retried["processed"] == 2 and retried["failed"] == 1
assert (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
"collection_version"
] == version
state["fields"] = field_records("TEST_UNKNOWN", 0)
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
await runner.execute(job.id)
done = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert done["status"] == "completed" and done["processed"] == 3 and done["failed"] == 0
async def test_cli_timeout_reuses_active_job_and_does_not_cancel(catalog, monkeypatch, capsys):
from app import cli
_, runner, _ = catalog
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
args = Namespace(
region="USA",
universe="TOP3000",
delay=1,
instrument_type="EQUITY",
resume_job=None,
wait_timeout=0.01,
)
assert await cli.catalog_sync_command(args) == 4
assert await cli.catalog_sync_command(args) == 4
async with runner.sessions() as db:
rows = list(await db.scalars(select(Job).where(Job.kind == "catalog_full_sync")))
assert len(rows) == 1 and rows[0].status == "queued" and not rows[0].cancel_requested
assert "等待超时" in capsys.readouterr().out
async def test_collection_version_checked_at_research_submission(catalog):
client = catalog[0]
collection = await copied(catalog)
response = await client.patch(f"{BASE}/{collection['id']}", json={"version": 1, "name": "changed"})
assert response.status_code == 200
body = {
"inline": {
"name": "prepared backtest",
"preparation_refs": [{"id": collection["id"], "version": 1}],
"candidates": [
{
"client_item_id": "one",
"expression": "rank(TEST_FIN_001)",
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
}
],
}
}
assert (await client.post("/api/v1/backtests/previews", json=body)).status_code == 409
body["inline"]["preparation_refs"][0]["version"] = 2
response = await client.post("/api/v1/backtests/previews", json=body)
assert response.status_code == 201, response.text
snapshot_id = response.json()["source"]["input_snapshot_ids"][0]
await client.delete(f"{BASE}/{collection['id']}?version=2")
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot_id}")).json()["fields"][1][
"dataset_id"
] == "TEST_FIN"
@pytest.mark.parametrize(
"status,code",
[
("completed", 0),
("completed_with_errors", 1),
("failed", 1),
("cancelled", 1),
("waiting_connection", 3),
("waiting_auth", 3),
],
)
async def test_cli_terminal_exit_codes(catalog, monkeypatch, status, code):
from app import cli
_, runner, _ = catalog
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
original = Catalog.create_job
async def terminal(self, body, **kwargs):
result = await original(self, body, **kwargs)
row = await self.db.get(Job, result.id)
row.status = status
return result
monkeypatch.setattr(Catalog, "create_job", terminal)
args = Namespace(
region="USA", universe="TOP3000", delay=1, instrument_type="EQUITY", resume_job=None, wait_timeout=1
)
assert await cli.catalog_sync_command(args) == code
async def test_full_sync_restart_keeps_page_and_auth_pauses_all_datasets(catalog, monkeypatch):
import asyncio
from app.worldquant import WqError
client, runner, _ = catalog
async with runner.sessions.begin() as db:
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
original = runner.client.catalog_page
interrupted = False
async def pages(scope, dataset_id, offset):
nonlocal interrupted
if dataset_id == "TEST_FIN" and offset == 100 and not interrupted:
interrupted = True
runner.stopping = True
raise asyncio.CancelledError()
if dataset_id == "TEST_NEWS":
raise WqError("synthetic network interruption", "network_error")
return await original(scope, dataset_id, offset)
monkeypatch.setattr(runner.client, "catalog_page", pages)
await runner.execute(job.id)
async with runner.sessions() as db:
batch = await db.scalar(
select(CatalogBatch).where(CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_FIN")
)
assert batch.offset == 100 and not batch.complete
assert (await db.get(Job, job.id)).status == "queued"
runner.stopping = False
await runner.execute(job.id)
row = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
assert row["status"] == "waiting_connection" and row["processed"] == 1
assert row["checkpoint"]["dataset_id"] == "TEST_NEWS"
assert row["checkpoint"]["datasets_completed"] == 1
async with runner.sessions() as db:
assert not await db.scalar(
select(CatalogBatch).where(
CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_UNKNOWN"
)
)
await client.post(f"/api/v1/sync-jobs/{job.id}/cancel")
assert (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()["status"] == "cancelled"
async def test_online_instrument_type_is_verified(catalog):
client, _, state = catalog
state["fields"][0]["instrumentType"] = "FUTURE"
assert (await client.get("/api/v1/catalog/worldquant/fields", params=SCOPE)).status_code == 502
async def test_mcp_collection_reads_and_versioned_submit_contract(catalog):
from types import SimpleNamespace
from app.research_access.contracts import PreparationRead, PreparationSearch, Submit
from app.research_access.service import ResearchAccess
client, runner, _ = catalog
collection = await copied(catalog)
async with runner.sessions.begin() as db:
access = ResearchAccess(
db, SimpleNamespace(token_id="test", admin_id=1), runner.client, "http://testserver"
)
result = await access.preparations(PreparationSearch(q=collection["name"]))
assert result["items"][0]["version"] == 1
detail = await access.preparation(PreparationRead(id=collection["id"], version=1, limit=1))
assert detail["fields"]["has_more"] and detail["fields"]["items"][0]["dataset_id"] == "TEST_FIN"
# Schema accepts versioned references and rejects accidental old dataset input payloads.
from tests.test_mcp import submission
payload = submission()
payload["preparation_refs"] = [{"id": collection["id"], "version": 1}]
assert Submit.model_validate(payload).preparation_refs[0].id == collection["id"]
async def test_retry_full_job_reuses_another_active_job(catalog):
from app.business import Business
_, runner, _ = catalog
body = CatalogJobInput(scope=Scope(**SCOPE))
async with runner.sessions.begin() as db:
old = await Catalog(db).create_job(body, full=True)
(await db.get(Job, old.id)).status = "failed"
async with runner.sessions.begin() as db:
new = await Catalog(db).create_job(body, full=True)
assert new.id != old.id
async with runner.sessions.begin() as db:
result = await Business(db).retry_job(old.id)
assert result["id"] == new.id
assert (await db.get(Job, old.id)).status == "failed"
async def test_retry_waiting_full_job_requeues_its_checkpoint(catalog):
from app.business import Business
_, runner, _ = catalog
async with runner.sessions.begin() as db:
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
row = await db.get(Job, job.id)
row.status, row.checkpoint = "waiting_connection", {"offset": 100}
async with runner.sessions.begin() as db:
result = await Business(db).retry_job(job.id)
assert result["status"] == "queued" and result["checkpoint"]["offset"] == 100
+12 -24
View File
@@ -11,7 +11,7 @@ from app.ai.tools import CAPABILITIES
from app.alphas import upsert_alpha from app.alphas import upsert_alpha
from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput
from app.business import Business 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_ai import configure, single_tool_factory
from tests.test_backtests import execute, setup, start from tests.test_backtests import execute, setup, start
from tests.test_catalog import SCOPE, prepare, sync from tests.test_catalog import SCOPE, prepare, sync
@@ -47,7 +47,7 @@ def construction(input_id):
return { return {
"name": "字段研究", "name": "字段研究",
"hypothesis": "显式字段排序", "hypothesis": "显式字段排序",
"template_input_id": input_id, "input_snapshot_id": input_id,
"candidates": [ "candidates": [
{ {
"client_item_id": "one", "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"] conversation = (await logged_in.post("/api/v1/ai/conversations")).json()["id"]
context = {"page": "datasets", "catalog_scope": SCOPE, "dataset_id": "TEST_FIN"} context = {"page": "datasets", "catalog_scope": SCOPE, "dataset_id": "TEST_FIN"}
if use_saved_input: 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) run = await ask(logged_in, conversation, "研究此输入" if use_saved_input else "自行选字段研究", context)
assert run["status"] == "waiting_approval", run assert run["status"] == "waiting_approval", run
assert not platform.posts 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["kind"] == "chatbox"
assert source["reference"] == conversation assert source["reference"] == conversation
assert source["research_id"] == run["id"] 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)" assert approval["preview"]["backtest"]["items"][0]["expression"] == "rank(TEST_FIN_001)"
async with app.state.sessions() as db: async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0 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": elif invalid == "unknown_type":
item["bindings"]["signal"] = {"field_id": "TEST_FIN_122", "field_type": "MATRIX"} item["bindings"]["signal"] = {"field_id": "TEST_FIN_122", "field_type": "MATRIX"}
else: else:
body["template_input_id"] = "missing" body["input_snapshot_id"] = "missing"
response = await logged_in.post("/api/v1/backtests/research-previews", json=body) 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 assert response.status_code == (404 if invalid == "missing_input" else 422), response.text
async with app.state.sessions() as db: 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["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["items"][-1]["field_type"] == "FUTURE_TYPE" and not page["has_more"]
assert page["_meta"]["source"] == "local_database" assert page["_meta"]["source"] == "local_database"
selected = await tool( subset = (await logged_in.post("/api/v1/data-preparations", json={"name": "one field", "scope": SCOPE,
"prepare_research_input", "fields": [{"scope": SCOPE, "dataset_id": "TEST_FIN", "field_id": "TEST_FIN_001", "source": "local",
{ "collection_version": fixed_input["fields"][0]["collection_version"]}]})).json()
"scope": SCOPE, selected = await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
)
assert selected["field_count"] == 1 assert selected["field_count"] == 1
bad = construction(selected["id"]) bad = construction(selected["id"])
bad["candidates"][0]["bindings"]["signal"]["field_id"] = "TEST_FIN_002" 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"]) "/api/v1/backtests/research-previews", json=construction(fixed_input["id"])
) )
assert response.status_code == 201, response.text 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: with pytest.raises(HTTPException) as exc:
await tool( await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
"prepare_research_input",
{
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": fixed_input["collection_version"],
"field_ids": ["TEST_FIN_001"],
},
)
assert exc.value.status_code == 409 assert exc.value.status_code == 409
async with app.state.sessions() as db: 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): async def test_multiple_origins_preserve_research_and_do_not_duplicate_alphas(app, logged_in):
+20 -8
View File
@@ -213,7 +213,7 @@ async def test_operator_annotation_and_refresh_preserve_local(app, logged_in, re
async def test_settings_variant_requires_all_fields_in_target(app, logged_in, research_input, catalog): 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.experiments import Experiments
from app.research.workspace_contracts import SettingVariants 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)) db.add(CatalogScope(key=target_key, scope=target_scope))
await db.flush() await db.flush()
db.add( db.add(
TemplateInput( ResearchInputSnapshot(
id="target", id="target",
scope_key=target_key, preparation_id="target-preparation",
dataset_id="TEST_FIN", preparation_version=1,
collection_version=research_input["collection_version"], content={**{k: v for k, v in research_input.items() if k not in ("id", "preparation_id", "preparation_version", "created_at")}, "scope": target_scope,
selection="explicit", "field_ids": ["TEST_FIN_001"], "field_types": {"TEST_FIN_001": "MATRIX"},
field_ids=["TEST_FIN_001"], "fields": [f for f in research_input["fields"] if f["id"] == "TEST_FIN_001"]},
field_types={"TEST_FIN_001": "MATRIX"},
) )
) )
await db.flush() await db.flush()
@@ -514,3 +513,16 @@ def test_real_seed_settings_preserve_execution_options_and_reject_unknowns():
assert snapshot["startDate"] == "2014-01-01" assert snapshot["startDate"] == "2014-01-01"
with pytest.raises(ValidationError): with pytest.raises(ValidationError):
seed_settings({**snapshot, "unknownOption": True}) 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"))
+27
View File
@@ -87,3 +87,30 @@ docker logs --tail=100 wq-alpha-production-web-1
采用 compound 经验 1 中的生产/开发隔离、显式外部网络、必填凭据、固定项目名、独立迁移和部署健康检查。未照搬旧项目端口、业务 Job 或卷权限初始化:本应用后端不写 named volume,已有镜像使用 UID 10001。 采用 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)。 凭据存放在 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 实际触发效果不属于模拟验收。
+5 -1
View File
@@ -51,6 +51,8 @@ python -m app.cli mcp-token-revoke TOKEN_ID
| authenticate_worldquant | `{action?:"connect"或"verify"}`;默认 connect,使用已保存凭据异步认证,返回 job_id;要求 research:refresh | | authenticate_worldquant | `{action?:"connect"或"verify"}`;默认 connect,使用已保存凭据异步认证,返回 job_id;要求 research:refresh |
| get_research_capabilities | `{}`,含单次候选上限及完整候选 schema | | get_research_capabilities | `{}`,含单次候选上限及完整候选 schema |
| create_research_template | `{template,hypothesis,source_item_ids,idempotency_key,reference?}`;保存调用方生成的完整模板,要求 research:write | | 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 查数据集,提供则查字段 | | search_catalog | `{filters:{region,universe,delay,...},dataset_id?}`;省略 dataset_id 查数据集,提供则查字段 |
| get_research_metadata | `{query:{kind,...}}`;kind 为 scopes/settings/operators/field_availability | | get_research_metadata | `{query:{kind,...}}`;kind 为 scopes/settings/operators/field_availability |
| refresh_research_data | `{query:{kind,...}}`;kind 为 catalog/operators/settings/field_availability/pnl | | 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 | | check_self_correlation | `{alpha_ids:[...]}`,1–100 个已导入 Alpha ID;异步返回 job_id |
| get_self_correlation | `{alpha_id}`;只读最新本地结果,含缓存和 stale 状态 | | get_self_correlation | `{alpha_id}`;只读最新本地结果,含缓存和 stale 状态 |
| search_backtests | 来源、reference、status、带时区起止时间、scope、q、候选精确匹配及分页 | | 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 | `{run_id,after?,event_limit?}`,after 为事件游标 |
| get_backtest_results | `{run_id,item_ids?,limit?,offset?}` | | get_backtest_results | `{run_id,item_ids?,limit?,offset?}` |
| get_backtest_artifact | `{item_id,kind,limit?,offset?,date_from?,date_to?}`,kind 为 snapshot/pnl | | 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}`。提交时核对版本、范围和字段,固定独立快照并保存到回测来源;空集合或版本冲突不会创建运行。后续编辑或删除集合不影响回测。无需旧输入草稿接口。
+15
View File
@@ -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 { useCallback, useEffect, useRef, useState } from "react";
import { import {
Badge, Badge,
@@ -390,6 +393,15 @@ export default function App() {
onTask={taskCreated} onTask={taskCreated}
/> />
</div> </div>
<div className="alpha-page-view" hidden={page !== "fields"}>
<FieldDirectory
active={page === "fields"}
revision={completedVersion}
/>
</div>
<div className="alpha-page-view" hidden={page !== "preparations"}>
<DataPreparationPage active={page === "preparations"} />
</div>
<div className="alpha-page-view" hidden={page !== "datasets"}> <div className="alpha-page-view" hidden={page !== "datasets"}>
<DatasetPage <DatasetPage
onContext={setDatasetContext} onContext={setDatasetContext}
@@ -507,6 +519,7 @@ export default function App() {
<ModelSettingsPanel /> <ModelSettingsPanel />
<WorkspacePreferences account={account} onChange={actionDone} /> <WorkspacePreferences account={account} onChange={actionDone} />
</SideSheet> </SideSheet>
<SnapshotDialog action={aiAction} />
<JobPanel <JobPanel
visible={showJobs && !(viewport < 1440 && chatOpen)} visible={showJobs && !(viewport < 1440 && chatOpen)}
chatOffset={chatOffset} chatOffset={chatOffset}
@@ -560,6 +573,8 @@ export default function App() {
: { page: "variants" as const }, : { page: "variants" as const },
alphas: alphaContext, alphas: alphaContext,
datasets: datasetContext, datasets: datasetContext,
fields: { page: "fields" as const },
preparations: { page: "preparations" as const },
backtests: backtestContext, backtests: backtestContext,
account: { page: "account" as const }, account: { page: "account" as const },
"mcp-keys": { page: "account" as const }, "mcp-keys": { page: "account" as const },
+3 -1
View File
@@ -17,6 +17,8 @@ export type PageContext = {
| "alphas" | "alphas"
| "account" | "account"
| "datasets" | "datasets"
| "fields"
| "preparations"
| "backtests" | "backtests"
| "operators" | "operators"
| "templates" | "templates"
@@ -36,7 +38,7 @@ export type PageContext = {
dataset_id?: string; dataset_id?: string;
field_id?: string; field_id?: string;
collection_version?: string; collection_version?: string;
template_input_id?: string; input_snapshot_id?: string;
unsaved_field_selection?: boolean; unsaved_field_selection?: boolean;
backtest_run_id?: string; backtest_run_id?: string;
backtest_preview_id?: string; backtest_preview_id?: string;
+6 -4
View File
@@ -12,7 +12,7 @@ export function contextReferences(context: PageContext): ResearchReference[] {
add("scope", "范围", `${scope.region}/${scope.universe}/D${scope.delay}`); add("scope", "范围", `${scope.region}/${scope.universe}/D${scope.delay}`);
add("dataset", "数据集", context.dataset_id); add("dataset", "数据集", context.dataset_id);
add("field", "字段", context.field_id); add("field", "字段", context.field_id);
add("input", "固定输入", context.template_input_id); add("input", "固定输入", context.input_snapshot_id);
add("version", "集合版本", context.collection_version); add("version", "集合版本", context.collection_version);
add("asset", "研究素材", context.research_asset_id); add("asset", "研究素材", context.research_asset_id);
add("experiment", "实验", context.research_experiment_id); add("experiment", "实验", context.research_experiment_id);
@@ -29,7 +29,7 @@ export function researchPrompts(context: PageContext): string[] {
case "datasets": case "datasets":
return context.unsaved_field_selection return context.unsaved_field_selection
? ["解释当前数据集的字段和适用场景"] ? ["解释当前数据集的字段和适用场景"]
: context.template_input_id : context.input_snapshot_id
? [ ? [
"分析这个固定输入中的字段,提出可验证的研究假设", "分析这个固定输入中的字段,提出可验证的研究假设",
"基于这个固定输入构建候选,并预览回测", "基于这个固定输入构建候选,并预览回测",
@@ -63,6 +63,8 @@ const contextLabels: Record<
PageContext["page"], PageContext["page"],
(context: PageContext) => string (context: PageContext) => string
> = { > = {
fields: () => "上下文:字段目录",
preparations: () => "上下文:数据准备",
home: () => "上下文:首页看板", home: () => "上下文:首页看板",
alphas: (context) => alphas: (context) =>
`上下文:${context.alpha_id ? `Alpha ${context.alpha_id}` : "Alpha 列表"}${context.selected_ids?.length ? ` · 已选 ${context.selected_ids.length} 条` : ""}`, `上下文:${context.alpha_id ? `Alpha ${context.alpha_id}` : "Alpha 列表"}${context.selected_ids?.length ? ` · 已选 ${context.selected_ids.length} 条` : ""}`,
@@ -78,7 +80,7 @@ const contextLabels: Record<
account: () => "上下文:个人信息页", account: () => "上下文:个人信息页",
backtests: () => "上下文:回测研究", backtests: () => "上下文:回测研究",
datasets: (context) => datasets: (context) =>
`上下文:${context.dataset_id ?? "数据目录"}${context.catalog_scope ? ` · ${context.catalog_scope.region}/${context.catalog_scope.universe}/D${context.catalog_scope.delay}` : ""}${context.template_input_id ? " · 固定研究输入" : context.unsaved_field_selection ? " · 请先保存字段选择" : ""}(不发送未保存备注)`, `上下文:${context.dataset_id ?? "数据目录"}${context.catalog_scope ? ` · ${context.catalog_scope.region}/${context.catalog_scope.universe}/D${context.catalog_scope.delay}` : ""}${context.input_snapshot_id ? " · 固定研究输入" : context.unsaved_field_selection ? " · 请先保存字段选择" : ""}(不发送未保存备注)`,
}; };
export function contextLabel(context: PageContext): string { export function contextLabel(context: PageContext): string {
@@ -103,7 +105,7 @@ const destinations: Record<UIAction["type"], Destination> = {
open_experiment: { page: "templates", chat: "responsive" }, open_experiment: { page: "templates", chat: "responsive" },
open_variant: { page: "variants", chat: "responsive" }, open_variant: { page: "variants", chat: "responsive" },
open_conversation: { chat: "open" }, open_conversation: { chat: "open" },
open_research_input: { page: "datasets", chat: "close" }, open_research_input: { page: "preparations", chat: "close" },
open_backtest: { page: "backtests", chat: "responsive" }, open_backtest: { page: "backtests", chat: "responsive" },
open_backtest_preview: { page: "backtests", chat: "responsive" }, open_backtest_preview: { page: "backtests", chat: "responsive" },
open_alpha: { page: "alphas", chat: "responsive" }, open_alpha: { page: "alphas", chat: "responsive" },
+1
View File
@@ -82,6 +82,7 @@ export const stateOptions = Object.entries(stateLabels).map(
([value, label]) => ({ value, label }), ([value, label]) => ({ value, label }),
); );
export const jobLabels: Record<string, string> = { export const jobLabels: Record<string, string> = {
catalog_full_sync: "全量同步数据目录与字段",
catalog_sync: "同步数据集目录", catalog_sync: "同步数据集目录",
field_sync: "同步数据字段", field_sync: "同步数据字段",
full_sync: "全量同步 Alpha", full_sync: "全量同步 Alpha",
+29 -1
View File
@@ -1,3 +1,8 @@
import {
ResearchDataInput,
researchSelection,
} from "../preparations/ResearchDataInput";
import type { InputSnapshot } from "../research/workspaceTypes";
import { useCallback, useEffect, useState } from "react"; import { useCallback, useEffect, useState } from "react";
import { import {
Banner, Banner,
@@ -70,6 +75,8 @@ export function BacktestPage({
const [editor, setEditor] = useState(false); const [editor, setEditor] = useState(false);
const [draft, setDraft] = useState<Draft | null>(null); const [draft, setDraft] = useState<Draft | null>(null);
const [name, setName] = useState(""); const [name, setName] = useState("");
const [inputIds, setInputIds] = useState<string[]>([]);
const [inputs, setInputs] = useState<InputSnapshot[]>([]);
const [source, setSource] = useState<Source>({ kind: "manual" }); const [source, setSource] = useState<Source>({ kind: "manual" });
const [text, setText] = useState(""); const [text, setText] = useState("");
const [mode, setMode] = useState("lines"); const [mode, setMode] = useState("lines");
@@ -215,6 +222,8 @@ export function BacktestPage({
setDraft(null); setDraft(null);
setName(""); setName("");
setSource({ kind: "manual" }); setSource({ kind: "manual" });
setInputIds([]);
setInputs([]);
setText(""); setText("");
setSettings(initialSettings); setSettings(initialSettings);
setMode("lines"); setMode("lines");
@@ -241,7 +250,7 @@ export function BacktestPage({
})); }));
if (!Array.isArray(candidates) || !candidates.length) if (!Array.isArray(candidates) || !candidates.length)
throw new Error("请提供非空候选集合"); throw new Error("请提供非空候选集合");
return { name, source, candidates }; return { name, source, candidates, ...researchSelection(inputIds, inputs) };
} }
async function loadDraft(id: string) { async function loadDraft(id: string) {
setSettingsDraft(null); setSettingsDraft(null);
@@ -249,6 +258,8 @@ export function BacktestPage({
setDraft(d); setDraft(d);
setName(d.name); setName(d.name);
setSource(d.source); setSource(d.source);
setInputIds(d.source.input_snapshot_ids ?? []);
setInputs([]);
setMode("json"); setMode("json");
setText(JSON.stringify(d.candidates, null, 2)); setText(JSON.stringify(d.candidates, null, 2));
setPreview(null); setPreview(null);
@@ -600,6 +611,23 @@ export function BacktestPage({
maxLength={200} maxLength={200}
/> />
</label> </label>
<ResearchDataInput
label="选择回测数据准备(可选)"
ids={inputIds}
inputs={inputs}
onChange={(ids, rows) => {
setInputIds(ids);
setInputs(rows);
const first = rows.find((r) => r.id === ids[0]);
if (first)
setSettings((old) => ({
...old,
region: first.scope.region,
universe: first.scope.universe,
delay: first.scope.delay,
}));
}}
/>
<div className="backtest-form-grid"> <div className="backtest-form-grid">
<label> <label>
来源 来源
+2 -1
View File
@@ -39,7 +39,8 @@ export type Source = {
kind: string; kind: string;
reference?: string | null; reference?: string | null;
batch_id?: string | null; batch_id?: string | null;
template_input_id?: string | null; input_snapshot_ids?: string[];
input_snapshot_id?: string | null;
research_id?: string | null; research_id?: string | null;
parent_run_id?: string | null; parent_run_id?: string | null;
hypothesis?: string | null; hypothesis?: string | null;
+7
View File
@@ -18,6 +18,13 @@ import type { WorkspacePage } from "../ai/workspace";
const navigation = [ const navigation = [
{ id: "home", label: "首页看板", icon: IconGridView, group: "概览" }, { id: "home", label: "首页看板", icon: IconGridView, group: "概览" },
{ id: "datasets", label: "数据目录", icon: IconList, group: "数据与素材" }, { id: "datasets", label: "数据目录", icon: IconList, group: "数据与素材" },
{ id: "fields", label: "字段目录", icon: IconList, group: "数据与素材" },
{
id: "preparations",
label: "数据准备",
icon: IconList,
group: "数据与素材",
},
{ id: "operators", label: "算子库", icon: IconList, group: "数据与素材" }, { id: "operators", label: "算子库", icon: IconList, group: "数据与素材" },
{ {
id: "templates", id: "templates",
+8
View File
@@ -93,6 +93,14 @@ export function JobPanel({
{job.checkpoint.dates_completed} / {job.checkpoint.dates_total} 天 {job.checkpoint.dates_completed} / {job.checkpoint.dates_total} 天
</p> </p>
)} )}
{job.kind === "catalog_full_sync" && (
<p>
数据集 {String(job.checkpoint?.dataset_id ?? "目录")} · 已完成{" "}
{String(job.checkpoint?.datasets_completed ?? 0)} /{" "}
{String(job.checkpoint?.datasets_total ?? "待统计")}
{!!job.error && <span> · {job.error}</span>}
</p>
)}
{job.kind === "pnl_backfill" && ( {job.kind === "pnl_backfill" && (
<p className="muted"> <p className="muted">
{job.total === 0 {job.total === 0
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,421 @@
import { useEffect, useState } from "react";
import {
Banner,
Button,
Input,
Modal,
Pagination,
Table,
TextArea,
Toast,
Tabs,
TabPane,
} from "@douyinfe/semi-ui-19";
import { api, patch, post, queryString, formatTime } from "../api";
import {
type Page,
type Preparation,
type FieldRecord,
defaultScope,
fieldReference,
scopeKey,
scopeLabel,
} from "./types";
import { ScopeControls } from "./ScopeControls";
import { FieldTable } from "./FieldTable";
import { FieldBrowser } from "./FieldDirectory";
import "./style.css";
export function DataPreparationPage({ active }: { active: boolean }) {
const [data, setData] = useState<Page<Preparation>>({
items: [],
total: 0,
limit: 25,
offset: 0,
}),
[q, setQ] = useState(""),
[page, setPage] = useState(1),
[rev, setRev] = useState(0);
const [selected, setSelected] = useState<Record<string, Preparation>>({}),
[scope, setScope] = useState(defaultScope),
[scopeFilter, setScopeFilter] = useState(false);
const [detail, setDetail] = useState<Preparation | null>(null),
[creating, setCreating] = useState(false),
[name, setName] = useState(""),
[note, setNote] = useState("");
const [fields, setFields] = useState<Page<FieldRecord>>({
items: [],
total: 0,
limit: 25,
offset: 0,
}),
[fieldPage, setFieldPage] = useState(1),
[fieldQ, setFieldQ] = useState("");
const [fieldSelection, setFieldSelection] = useState<string[]>([]),
[adding, setAdding] = useState(false),
[tab, setTab] = useState("local");
const [error, setError] = useState(""),
[busy, setBusy] = useState(false),
[deleting, setDeleting] = useState<Preparation[]>([]);
useEffect(() => {
if (!active) return;
const c = new AbortController();
api<Page<Preparation>>(
`/data-preparations?${queryString({ q, offset: (page - 1) * 25, scope_key: scopeFilter ? scopeKey(scope) : undefined })}`,
{ signal: c.signal },
)
.then((r) => {
setData(r);
setPage((p) => Math.min(p, Math.max(1, Math.ceil(r.total / 25))));
})
.catch((e) => {
if (!c.signal.aborted) setError(e.message);
});
return () => c.abort();
}, [active, q, page, rev, scopeFilter, scopeKey(scope)]);
useEffect(() => {
if (!detail) return;
const c = new AbortController();
api<Page<FieldRecord>>(
`/data-preparations/${detail.id}/fields?${queryString({ q: fieldQ, offset: (fieldPage - 1) * 25 })}`,
{ signal: c.signal },
)
.then(setFields)
.catch((e) => {
if (!c.signal.aborted) setError(e.message);
});
return () => c.abort();
}, [detail, fieldQ, fieldPage]);
async function work(fn: () => Promise<void>) {
setBusy(true);
setError("");
try {
await fn();
setRev((v) => v + 1);
} catch (e) {
setError((e as Error).message);
} finally {
setBusy(false);
}
}
function open(row: Preparation) {
setDetail(row);
setName(row.name);
setNote(row.note);
setFieldQ("");
setFieldPage(1);
setFieldSelection([]);
setError("");
}
useEffect(() => {
if (!active) return;
const id = new URLSearchParams(location.hash.split("?")[1] ?? "").get("id");
if (id)
void work(async () =>
open(
await api<Preparation>(
`/data-preparations/${encodeURIComponent(id)}`,
),
),
);
}, [active]);
return (
<section className="preparation-page">
{error && <Banner type="danger" description={error} />}
<div className="preparation-tools">
<Input
aria-label="搜索准备集合"
placeholder="搜索名称或备注"
value={q}
onChange={(v) => {
setQ(v);
setPage(1);
}}
/>
<Button
theme="solid"
onClick={() => {
setCreating(true);
setName("新数据准备集合");
setNote("");
}}
>
新建集合
</Button>
<Button
onClick={() => {
setScopeFilter((v) => !v);
setPage(1);
}}
>
{scopeFilter ? "全部范围" : "按范围筛选"}
</Button>
<Button
type="danger"
disabled={!Object.keys(selected).length}
onClick={() => setDeleting(Object.values(selected))}
>
删除所选
</Button>
</div>
{scopeFilter && (
<ScopeControls
local
value={scope}
onChange={(s) => {
setScope(s);
setPage(1);
setSelected({});
}}
/>
)}
<Table<Preparation>
rowKey="id"
size="small"
dataSource={data.items}
pagination={false}
scroll={{ x: 1000, y: "100%" }}
rowSelection={{
selectedRowKeys: Object.keys(selected),
onChange: (keys) =>
setSelected((old) => {
const n: Record<string, Preparation> = {};
for (const k of keys ?? []) {
const r = data.items.find((x) => x.id === k) ?? old[k];
if (r) n[k] = r;
}
return n;
}),
}}
columns={[
{
title: "集合名称",
render: (_, r) => (
<Button theme="borderless" onClick={() => open(r)}>
{r.name}
</Button>
),
},
{ title: "范围", render: (_, r) => scopeLabel(r.scope) },
{ title: "字段数", dataIndex: "field_count" },
{ title: "数据集数", dataIndex: "dataset_count" },
{ title: "更新时间", render: (_, r) => formatTime(r.updated_at) },
{
title: "操作",
render: (_, r) => (
<div className="preparation-tools">
<Button theme="borderless" onClick={() => open(r)}>
查看
</Button>
<Button
theme="borderless"
disabled={busy}
onClick={() =>
void work(async () => {
const copy = await post<Preparation>(
`/data-preparations/${r.id}/copy`,
{ version: r.version },
);
open(copy);
})
}
>
复制
</Button>
<Button
theme="borderless"
type="danger"
onClick={() => setDeleting([r])}
>
删除
</Button>
</div>
),
},
]}
/>
<Pagination
currentPage={page}
total={data.total}
pageSize={25}
onPageChange={setPage}
/>
<Modal
visible={creating && active}
title="新建数据准备集合"
closeOnEsc
onCancel={() => setCreating(false)}
confirmLoading={busy}
okButtonProps={{ disabled: !name.trim() || busy }}
onOk={() =>
work(async () => {
const row = await post<Preparation>("/data-preparations", {
name,
note,
scope,
});
setCreating(false);
open(row);
})
}
>
{error && <Banner type="danger" description={error} />}
<Input aria-label="集合名称" value={name} onChange={setName} />
<ScopeControls value={scope} onChange={setScope} />
<TextArea aria-label="集合备注" value={note} onChange={setNote} />
</Modal>
<Modal
visible={!!detail && active && !adding && !deleting.length}
title="数据准备详情"
width="min(1200px, 96vw)"
closeOnEsc
footer={null}
onCancel={() => setDetail(null)}
>
{detail && (
<div className="preparation-detail">
{error && <Banner type="danger" description={error} />}
<div>
{scopeLabel(detail.scope)} · v{detail.version} ·{" "}
{detail.field_count} 字段 / {detail.dataset_count} 数据集
</div>
<label>
集合名称
<Input aria-label="集合名称" value={name} onChange={setName} />
</label>
<label>
备注
<TextArea aria-label="集合备注" value={note} onChange={setNote} />
</label>
<div className="preparation-tools">
<Button
disabled={busy || !name.trim()}
onClick={() =>
void work(async () =>
setDetail(
await patch(`/data-preparations/${detail.id}`, {
version: detail.version,
name,
note,
}),
),
)
}
>
保存名称和备注
</Button>
<Button
disabled={busy}
onClick={() =>
void work(async () =>
open(await api(`/data-preparations/${detail.id}`)),
)
}
>
重新载入
</Button>
<Button onClick={() => setAdding(true)}>添加字段</Button>
<Button
disabled={busy || !fieldSelection.length}
onClick={() =>
void work(async () => {
setDetail(
await patch(`/data-preparations/${detail.id}/fields`, {
version: detail.version,
remove_ids: fieldSelection.map((k) => k.split("|")[1]),
}),
);
setFieldSelection([]);
})
}
>
移除所选字段
</Button>
<Input
aria-label="集合内搜索字段"
value={fieldQ}
placeholder="搜索字段名称、ID、描述"
onChange={(v) => {
setFieldQ(v);
setFieldPage(1);
}}
/>
</div>
<FieldTable
items={fields.items}
selected={fieldSelection}
onSelect={setFieldSelection}
/>
<Pagination
currentPage={fieldPage}
total={fields.total}
pageSize={25}
onPageChange={setFieldPage}
/>
</div>
)}
</Modal>
<Modal
visible={adding && active}
title="添加同范围字段"
width="min(1250px,96vw)"
footer={null}
closeOnEsc
onCancel={() => setAdding(false)}
>
{error && <Banner type="danger" description={error} />}
<Tabs activeKey={tab} onChange={setTab} tabPaneMotion={false}>
{(["local", "worldquant"] as const).map((source) => (
<TabPane
key={source}
itemKey={source}
tab={source === "local" ? "本地同步" : "worldquant接口"}
>
{detail && (
<FieldBrowser
active={adding && tab === source}
source={source}
fixedScope={detail.scope}
onAdd={(chosen) =>
void work(async () => {
setDetail(
await patch(`/data-preparations/${detail.id}/fields`, {
version: detail.version,
fields: chosen.map(fieldReference),
}),
);
setAdding(false);
})
}
/>
)}
</TabPane>
))}
</Tabs>
</Modal>
<Modal
visible={!!deleting.length}
title="删除数据准备集合"
closeOnEsc
onCancel={() => setDeleting([])}
confirmLoading={busy}
okType="danger"
okText="删除"
onOk={() =>
work(async () => {
await post("/data-preparations/batch-delete", {
items: deleting.map((r) => ({ id: r.id, version: r.version })),
});
if (deleting.some((r) => r.id === detail?.id)) setDetail(null);
setDeleting([]);
setSelected({});
Toast.success("集合已删除,已有研究快照保留");
})
}
>
{error && <Banner type="danger" description={error} />}删除{" "}
{deleting.length} 个集合:{deleting.map((r) => r.name).join("、")}
。已创建的研究保留使用时的快照。
</Modal>
</section>
);
}
@@ -0,0 +1,215 @@
import { useEffect, useState } from "react";
import {
Banner,
Button,
Input,
Modal,
Pagination,
Table,
} from "@douyinfe/semi-ui-19";
import { api, queryString } from "../api";
import { FieldTable } from "./FieldTable";
import {
type FieldRecord,
type Page,
type Preparation,
type PreparationRef,
type Scope,
scopeKey,
scopeLabel,
} from "./types";
export function DataPreparationPicker({
visible,
onCancel,
onConfirm,
scope,
multiple = true,
allowEmpty = false,
}: {
visible: boolean;
onCancel: () => void;
onConfirm: (
refs: PreparationRef[],
rows: Preparation[],
) => void | Promise<void>;
scope?: Scope;
multiple?: boolean;
allowEmpty?: boolean;
}) {
const [q, setQ] = useState(""),
[page, setPage] = useState(1),
[data, setData] = useState<Page<Preparation>>({
items: [],
total: 0,
limit: 25,
offset: 0,
});
const [selected, setSelected] = useState<Record<string, Preparation>>({}),
[error, setError] = useState(""),
[busy, setBusy] = useState(false);
const [preview, setPreview] = useState<Preparation | null>(null),
[fields, setFields] = useState<FieldRecord[]>([]),
[fieldPage, setFieldPage] = useState(1);
useEffect(() => {
if (visible) {
setSelected({});
setPreview(null);
setError("");
}
}, [visible, scope && scopeKey(scope)]);
useEffect(() => {
if (!visible) return;
const c = new AbortController();
setBusy(true);
api<Page<Preparation>>(
`/data-preparations?${queryString({ q, limit: 25, offset: (page - 1) * 25, scope_key: scope ? scopeKey(scope) : undefined })}`,
{ signal: c.signal },
)
.then(setData)
.catch((e) => {
if (!c.signal.aborted) setError(e.message);
})
.finally(() => {
if (!c.signal.aborted) setBusy(false);
});
return () => c.abort();
}, [visible, q, page, scope && scopeKey(scope)]);
useEffect(() => {
if (!preview) return;
const c = new AbortController();
setFields([]);
api<Page<FieldRecord>>(
`/data-preparations/${preview.id}/fields?limit=25&offset=${(fieldPage - 1) * 25}`,
{ signal: c.signal },
)
.then((r) => setFields(r.items))
.catch((e) => {
if (!c.signal.aborted) setError(e.message);
});
return () => c.abort();
}, [preview, fieldPage]);
return (
<Modal
visible={visible}
title="选择数据准备集合"
width="min(1080px, 96vw)"
closeOnEsc
onCancel={onCancel}
confirmLoading={busy}
okText="选择"
okButtonProps={{
"aria-label": "选择",
disabled: !Object.keys(selected).length || busy,
}}
onOk={async () => {
setBusy(true);
try {
const rows = Object.values(selected);
await onConfirm(
rows.map((r) => ({ id: r.id, version: r.version })),
rows,
);
} catch (e) {
setError((e as Error).message);
} finally {
setBusy(false);
}
}}
>
{error && <Banner type="danger" description={error} />}
<div className="preparation-tools">
<Input
aria-label="搜索数据准备集合"
placeholder="搜索名称或备注"
value={q}
onChange={(v) => {
setQ(v);
setPage(1);
}}
/>
<span>
{scope ? scopeLabel(scope) : "各集合保持各自范围"} · 已选{" "}
{Object.keys(selected).length}
</span>
</div>
<Table<Preparation>
rowKey="id"
dataSource={data.items}
size="small"
pagination={false}
loading={busy}
rowSelection={
multiple
? {
selectedRowKeys: Object.keys(selected),
getCheckboxProps: (r) => ({
disabled: !allowEmpty && !r.field_count,
}),
onChange: (keys) =>
setSelected((old) => {
const next: Record<string, Preparation> = {};
for (const k of keys ?? []) {
const row = data.items.find((r) => r.id === k) ?? old[k];
if (row) next[k] = row;
}
return next;
}),
}
: undefined
}
columns={[
{
title: "集合名称",
render: (_, r) => (
<Button
theme="borderless"
onClick={() => {
setPreview(r);
setFieldPage(1);
}}
>
{r.name}
</Button>
),
},
{ title: "范围", render: (_, r) => scopeLabel(r.scope) },
{ title: "字段数", dataIndex: "field_count" },
{ title: "数据集数", dataIndex: "dataset_count" },
...(!multiple
? [
{
title: "操作",
render: (_: unknown, r: Preparation) => (
<Button
disabled={!allowEmpty && !r.field_count}
onClick={() => setSelected({ [r.id]: r })}
>
{selected[r.id] ? "已选择" : "选择"}
</Button>
),
},
]
: []),
]}
/>
<Pagination
currentPage={page}
total={data.total}
pageSize={25}
onPageChange={setPage}
/>
{preview && (
<section>
<h4>{preview.name} · 字段预览</h4>
<FieldTable items={fields} />
<Pagination
currentPage={fieldPage}
total={preview.field_count}
pageSize={25}
onPageChange={setFieldPage}
/>
</section>
)}
</Modal>
);
}
@@ -0,0 +1,409 @@
import { useEffect, useState } from "react";
import {
Banner,
Button,
Input,
Modal,
Pagination,
TabPane,
Tabs,
Toast,
} from "@douyinfe/semi-ui-19";
import { api, patch, post, queryString } from "../api";
import { ResearchSelect } from "../research/ResearchSelect";
import { ScopeControls } from "./ScopeControls";
import { DataPreparationPicker } from "./DataPreparationPicker";
import { FieldTable } from "./FieldTable";
import {
defaultScope,
fieldReference,
scopeKey,
scopeLabel,
type FieldRecord,
type Page,
type Preparation,
type Scope,
} from "./types";
import "./style.css";
export function AddFieldsDialog({
fields,
onClose,
onDone,
}: {
fields: FieldRecord[];
onClose: () => void;
onDone?: () => void;
}) {
const [name, setName] = useState("新数据准备集合"),
[picker, setPicker] = useState(false),
[error, setError] = useState(""),
[busy, setBusy] = useState(false);
async function save(target?: Preparation) {
setBusy(true);
try {
const refs = fields.map(fieldReference);
const result = target
? await patch<Preparation>(`/data-preparations/${target.id}/fields`, {
version: target.version,
fields: refs,
})
: await post<Preparation>("/data-preparations", {
name,
scope: fields[0].scope,
fields: refs,
});
Toast.success(`已加入 ${result.name}`);
onDone?.();
onClose();
} catch (e) {
setError((e as Error).message);
} finally {
setBusy(false);
}
}
return (
<>
<Modal
visible={!picker}
title="加入数据准备"
closeOnEsc
onCancel={onClose}
onOk={() => save()}
okText="新建集合并加入"
confirmLoading={busy}
okButtonProps={{
"aria-label": "新建集合并加入",
disabled: !name.trim() || busy,
}}
>
{error && <Banner type="danger" description={error} />}
<p>
{fields.length} 个字段 · {scopeLabel(fields[0].scope)}
</p>
<Input aria-label="新集合名称" value={name} onChange={setName} />
<Button theme="borderless" onClick={() => setPicker(true)}>
追加到现有集合
</Button>
</Modal>
<DataPreparationPicker
visible={picker}
onCancel={() => setPicker(false)}
scope={fields[0].scope}
multiple={false}
allowEmpty
onConfirm={async (_, rows) => {
await save(rows[0]);
setPicker(false);
}}
/>
</>
);
}
export function FieldBrowser({
active = true,
source = "local",
fixedScope,
onAdd,
revision = "",
}: {
active?: boolean;
source?: "local" | "worldquant";
fixedScope?: Scope;
onAdd?: (fields: FieldRecord[]) => void;
revision?: string;
}) {
const [scope, setScope] = useState<Scope>(fixedScope ?? defaultScope),
[draft, setDraft] = useState<Record<string, string>>({}),
[filters, setFilters] = useState<Record<string, string>>({});
const [page, setPage] = useState(1),
[size, setSize] = useState(25),
[refresh, setRefresh] = useState(0),
[advanced, setAdvanced] = useState(false);
const [data, setData] = useState<Page<FieldRecord>>({
items: [],
total: 0,
limit: 25,
offset: 0,
}),
[selected, setSelected] = useState<Record<string, FieldRecord>>({});
const [error, setError] = useState(""),
[loading, setLoading] = useState(false),
[adding, setAdding] = useState(false);
const actual = fixedScope ?? scope,
key = scopeKey(actual);
useEffect(() => {
setSelected({});
setData({ items: [], total: 0, limit: 25, offset: 0 });
setPage(1);
}, [key]);
useEffect(() => {
if (!active) return;
const c = new AbortController();
setLoading(true);
setError("");
api<Page<FieldRecord>>(
`/catalog/${source === "local" ? "fields" : "worldquant/fields"}?${queryString({ ...actual, ...filters, limit: size, offset: (page - 1) * size })}`,
{ signal: c.signal },
)
.then((r) => {
setData(r);
setSelected((old) => {
const next = { ...old };
for (const f of r.items) {
const k = `${f.dataset_id}|${f.id}`;
if (next[k]) next[k] = f;
}
return next;
});
})
.catch((e) => {
if (!c.signal.aborted) setError(e.message);
})
.finally(() => {
if (!c.signal.aborted) setLoading(false);
});
return () => c.abort();
}, [active, source, key, filters, page, size, refresh, revision]);
const update = (k: string, v: string) =>
setDraft((old) => ({ ...old, [k]: v }));
return (
<div className="preparation-fields">
{!fixedScope && (
<ScopeControls
value={scope}
onChange={setScope}
local={source === "local"}
active={active}
revision={revision}
/>
)}
<div className="preparation-tools">
<Input
aria-label="搜索字段"
placeholder="字段 ID、名称或描述"
value={draft.q ?? ""}
onChange={(v) => update("q", v)}
/>
<Input
aria-label="字段数据集"
placeholder="数据集 ID"
value={draft.dataset_id ?? ""}
onChange={(v) => update("dataset_id", v)}
/>
<ResearchSelect
label="字段类型"
placeholder="全部类型"
showClear
value={draft.field_type || undefined}
optionList={["MATRIX", "VECTOR", "GROUP"].map((v) => ({
label: v,
value: v,
}))}
onChange={(v) => update("field_type", String(v ?? ""))}
/>
<Button
onClick={() => {
setFilters({
...draft,
synced_from: draft.synced_from
? new Date(draft.synced_from).toISOString()
: "",
synced_to: draft.synced_to
? new Date(draft.synced_to).toISOString()
: "",
});
setPage(1);
setRefresh((v) => v + 1);
}}
>
查询
</Button>
<Button onClick={() => setAdvanced((v) => !v)}>更多筛选</Button>
<Button
onClick={() => {
setDraft({});
setFilters({});
setPage(1);
}}
>
重置
</Button>
</div>
{advanced && (
<div className="preparation-filters">
{source === "local" &&
["category", "subcategory"].map((k, i) => (
<label key={k}>
{i ? "子分类" : "分类"}
<Input
aria-label={i ? "字段子分类" : "字段分类"}
value={draft[k] ?? ""}
onChange={(v) => update(k, v)}
/>
</label>
))}
{[
["coverage", "覆盖率(0–1)"],
["user_count", "用户数"],
["alpha_count", "Alpha 数"],
].flatMap(([k, label]) =>
["min", "max"].map((suffix) => (
<label key={k + suffix}>
{label}{" "}
{source === "worldquant"
? suffix === "min"
? ">"
: "<"
: suffix === "min"
? "≥"
: "≤"}
<Input
type="number"
aria-label={`${label}${suffix === "min" ? "下限" : "上限"}`}
value={draft[k + "_" + suffix] ?? ""}
onChange={(v) => update(k + "_" + suffix, v)}
/>
</label>
)),
)}
{source === "local" && (
<>
<label>
同步起始时间
<input
type="datetime-local"
value={draft.synced_from ?? ""}
onChange={(e) => update("synced_from", e.target.value)}
/>
</label>
<label>
同步结束时间
<input
type="datetime-local"
value={draft.synced_to ?? ""}
onChange={(e) => update("synced_to", e.target.value)}
/>
</label>
<ResearchSelect
label="字段排序"
value={draft.sort ?? "id"}
optionList={[
"id",
"name",
"dataset_id",
"coverage",
"user_count",
"alpha_count",
"synced_at",
].map((v) => ({ value: v, label: v }))}
onChange={(v) => update("sort", String(v))}
/>
<ResearchSelect
label="排序方向"
value={draft.direction ?? "asc"}
optionList={[
{ value: "asc", label: "升序" },
{ value: "desc", label: "降序" },
]}
onChange={(v) => update("direction", String(v))}
/>
</>
)}
</div>
)}
{error && <Banner type="danger" description={error} />}
<div className="preparation-tools">
<Button
theme="solid"
disabled={!Object.keys(selected).length || loading || !!error}
onClick={() =>
onAdd ? onAdd(Object.values(selected)) : setAdding(true)
}
>
加入数据准备
</Button>
<span>已选 {Object.keys(selected).length} 个字段 · 表头选择当前页</span>
{!!Object.keys(selected).length && (
<Button onClick={() => setSelected({})}>清空选择</Button>
)}
</div>
<FieldTable
items={data.items}
loading={loading}
selected={Object.keys(selected)}
onSelect={(ids) =>
setSelected((old) => {
const next: Record<string, FieldRecord> = {};
for (const id of ids) {
const f =
data.items.find((r) => `${r.dataset_id}|${r.id}` === id) ??
old[id];
if (f) next[id] = f;
}
return next;
})
}
/>
<div className="preparation-tools">
<span>
{data.total_known === false
? "平台未提供总数"
: `共 ${data.total} 个字段`}
</span>
<Pagination
currentPage={page}
pageSize={size}
total={data.total}
onPageChange={setPage}
/>
<ResearchSelect
label="字段每页条数"
value={size}
optionList={[25, 50, 100].map((v) => ({
value: v,
label: `${v} 条/页`,
}))}
onChange={(v) => {
setSize(Number(v));
setPage(1);
}}
/>
</div>
{adding && (
<AddFieldsDialog
fields={Object.values(selected)}
onClose={() => setAdding(false)}
onDone={() => setSelected({})}
/>
)}
</div>
);
}
export function FieldDirectory({
active,
revision,
}: {
active: boolean;
revision?: string;
}) {
const [tab, setTab] = useState("worldquant");
return (
<section className="preparation-page">
<Tabs activeKey={tab} onChange={setTab} keepDOM tabPaneMotion={false}>
<TabPane tab="worldquant接口" itemKey="worldquant">
<FieldBrowser
active={active && tab === "worldquant"}
source="worldquant"
revision={revision}
/>
</TabPane>
<TabPane tab="本地同步" itemKey="local">
<FieldBrowser
active={active && tab === "local"}
revision={revision}
/>
</TabPane>
</Tabs>
</section>
);
}
+94
View File
@@ -0,0 +1,94 @@
import { Button, Table } from "@douyinfe/semi-ui-19";
import type { ColumnProps } from "@douyinfe/semi-ui-19/lib/es/table/interface";
import { displayValue, formatTime } from "../api";
import type { FieldRecord } from "./types";
export function FieldTable({
items,
loading = false,
selected,
onSelect,
onOpen,
}: {
items: FieldRecord[];
loading?: boolean;
onOpen?: (field: FieldRecord) => void;
selected?: string[];
onSelect?: (ids: string[]) => void;
}) {
const columns: ColumnProps<FieldRecord>[] = [
{
title: "字段",
dataIndex: "id",
width: 220,
render: (_, r) =>
onOpen ? (
<Button theme="borderless" onClick={() => onOpen(r)}>
{r.id}
</Button>
) : (
r.id
),
},
{
title: "数据集",
width: 200,
render: (_, r) => <span title={r.dataset_name}>{r.dataset_id}</span>,
},
{
title: "字段描述",
dataIndex: "description",
width: 340,
render: displayValue,
},
{
title: "类型",
dataIndex: "field_type",
width: 100,
render: displayValue,
},
{
title: "覆盖率",
width: 100,
render: (_, r) =>
r.coverage == null ? "未提供" : `${(r.coverage * 100).toFixed(1)}%`,
},
{
title: "用户数",
dataIndex: "user_count",
width: 90,
render: displayValue,
},
{
title: "Alpha 数",
dataIndex: "alpha_count",
width: 100,
render: displayValue,
},
{
title: "同步时间",
width: 170,
render: (_, r) => (r.synced_at ? formatTime(r.synced_at) : "在线查询"),
},
];
return (
<Table<FieldRecord>
size="small"
className="preparation-table"
rowKey={(r) => `${r?.dataset_id}|${r?.id}`}
dataSource={items}
columns={columns}
loading={loading}
pagination={false}
scroll={{ x: 1420, y: 400 }}
rowSelection={
onSelect
? {
selectedRowKeys: selected ?? [],
onChange: (keys) => onSelect((keys ?? []).map(String)),
}
: undefined
}
empty="暂无匹配字段"
/>
);
}
@@ -0,0 +1,103 @@
import { useEffect, useState } from "react";
import { Banner, Button } from "@douyinfe/semi-ui-19";
import { api } from "../api";
import type { InputSnapshot } from "../research/workspaceTypes";
import { DataPreparationPicker } from "./DataPreparationPicker";
import { scopeLabel, type Scope } from "./types";
import "./style.css";
export function researchSelection(ids: string[], rows: InputSnapshot[]) {
const selected = rows.filter((r) => ids.includes(r.id));
return {
input_ids: ids.filter(
(id) => !selected.find((r) => r.id === id)?.preparation_ref,
),
preparation_refs: selected.flatMap((r) =>
r.preparation_ref ? [r.preparation_ref] : [],
),
};
}
export function ResearchDataInput({
ids,
inputs,
onChange,
scope,
multiple = true,
label = "选择数据准备",
}: {
ids: string[];
inputs: InputSnapshot[];
onChange: (ids: string[], inputs: InputSnapshot[]) => void;
scope?: Scope;
multiple?: boolean;
label?: string;
}) {
const [open, setOpen] = useState(false),
[error, setError] = useState("");
useEffect(() => {
const missing = ids.filter((id) => !inputs.some((r) => r.id === id));
if (!missing.length) return;
let live = true;
Promise.all(
missing.map((id) =>
api<InputSnapshot>(
`/research/input-snapshots/${encodeURIComponent(id)}`,
),
),
)
.then((rows) => {
if (live) onChange(ids, [...inputs, ...rows]);
})
.catch((e) => {
if (live) setError(e.message);
});
return () => {
live = false;
};
}, [ids.join("|"), inputs]);
return (
<div>
<div className="preparation-tools">
<Button aria-label={label} onClick={() => setOpen(true)}>
{label}
</Button>
{!!ids.length && (
<Button onClick={() => onChange([], inputs)}>清空选择</Button>
)}
</div>
{error && <Banner type="danger" description={error} />}
{inputs
.filter((r) => ids.includes(r.id))
.map((r) => (
<p key={r.id}>
{r.name} · {scopeLabel(r.scope)} · {r.field_ids.length} 字段
{r.preparation_ref
? ` · v${r.preparation_ref.version}`
: " · 研究快照"}
</p>
))}
<DataPreparationPicker
visible={open}
onCancel={() => setOpen(false)}
scope={scope}
multiple={multiple}
onConfirm={async (refs) => {
const rows = await Promise.all(
refs.map((r) =>
api<InputSnapshot>(
`/data-preparations/${r.id}/selection?version=${r.version}`,
),
),
);
onChange(
rows.map((r) => r.id),
[
...inputs.filter((old) => !rows.some((r) => r.id === old.id)),
...rows,
],
);
setOpen(false);
}}
/>
</div>
);
}
+132
View File
@@ -0,0 +1,132 @@
import { useEffect, useState } from "react";
import { Banner, Button } from "@douyinfe/semi-ui-19";
import { api } from "../api";
import { ResearchSelect } from "../research/ResearchSelect";
import { type Scope } from "./types";
type Option = Omit<Scope, "universe"> & { universes: string[] };
export function ScopeControls({
value,
onChange,
local = false,
active = true,
revision: sourceRevision = "",
}: {
value: Scope;
onChange: (s: Scope) => void;
local?: boolean;
active?: boolean;
revision?: string;
}) {
const [rows, setRows] = useState<Option[]>([]),
[error, setError] = useState(""),
[revision, setRevision] = useState(0);
useEffect(() => {
if (!active) return;
const c = new AbortController();
api<{ instrument_options: Option[] }>(
local ? "/catalog/local-scopes" : "/catalog/scopes",
{ signal: c.signal },
)
.then((r) => {
setRows(r.instrument_options);
setError("");
})
.catch((e) => {
if (!c.signal.aborted) setError(e.message);
});
return () => c.abort();
}, [local, revision, sourceRevision, active]);
function update(next: Scope) {
const row =
rows.find(
(r) =>
r.instrument_type === next.instrument_type &&
r.region === next.region &&
r.delay === next.delay,
) ??
rows.find(
(r) =>
r.instrument_type === next.instrument_type &&
r.region === next.region,
) ??
rows.find((r) => r.instrument_type === next.instrument_type) ??
rows[0];
onChange(
row
? {
...next,
instrument_type: row.instrument_type,
region: row.region,
delay: row.delay,
universe: row.universes.includes(next.universe)
? next.universe
: row.universes[0],
}
: next,
);
}
const option = rows.find(
(r) =>
r.instrument_type === value.instrument_type &&
r.region === value.region &&
r.delay === value.delay,
);
return (
<>
<div className="preparation-tools">
{(["instrument_type", "region", "universe", "delay"] as const).map(
(key) => {
const values =
key === "universe"
? (option?.universes ?? [value.universe])
: rows
.filter(
(r) =>
key === "instrument_type" ||
(r.instrument_type === value.instrument_type &&
(key !== "delay" || r.region === value.region)),
)
.map((r) => r[key]);
return (
<ResearchSelect
key={key}
label={
{
instrument_type: "InstrumentType",
region: "Region",
universe: "Universe",
delay: "Delay",
}[key]
}
value={value[key]}
optionList={[...new Set([value[key], ...values])].map((v) => ({
value: v,
label: key === "delay" ? `Delay ${v}` : String(v),
}))}
onChange={(v) =>
update({
...value,
[key]: key === "delay" ? Number(v) : String(v),
})
}
/>
);
},
)}
</div>
{error && (
<Banner
type="warning"
description={
<>
{error}{" "}
<Button onClick={() => setRevision((v) => v + 1)}>
重试范围选项
</Button>
</>
}
/>
)}
</>
);
}
@@ -0,0 +1,47 @@
import { useEffect, useState } from "react";
import { Banner, Modal } from "@douyinfe/semi-ui-19";
import { api } from "../api";
import type { UIAction } from "../ai/types";
import type { InputSnapshot } from "../research/workspaceTypes";
import { FieldTable } from "./FieldTable";
import { scopeLabel } from "./types";
export function SnapshotDialog({ action }: { action: UIAction | null }) {
const [snapshot, setSnapshot] = useState<InputSnapshot | null>(null),
[error, setError] = useState(""),
[open, setOpen] = useState(false);
useEffect(() => {
if (action?.type !== "open_research_input") return;
const c = new AbortController();
setOpen(true);
setSnapshot(null);
setError("");
api<InputSnapshot>(`/research/input-snapshots/${action.input_id}`, {
signal: c.signal,
})
.then(setSnapshot)
.catch((e) => {
if (!c.signal.aborted) setError(e.message);
});
return () => c.abort();
}, [action]);
return (
<Modal
visible={open}
title="研究输入快照"
width="min(1200px,96vw)"
closeOnEsc
footer={null}
onCancel={() => setOpen(false)}
>
{error && <Banner type="danger" description={error} />}
{snapshot && (
<>
<p>
{snapshot.name} · {scopeLabel(snapshot.scope)}
</p>
<FieldTable items={snapshot.fields} />
</>
)}
</Modal>
);
}
+68
View File
@@ -0,0 +1,68 @@
.preparation-page {
display: flex;
flex-direction: column;
gap: 12px;
height: 100%;
min-height: 0;
padding: 16px;
overflow: auto;
background: var(--semi-color-bg-0);
}
.preparation-tools {
display: flex;
align-items: center;
gap: 8px;
flex-wrap: wrap;
margin: 8px 0;
}
.preparation-tools > .semi-input-wrapper {
width: 240px;
}
.preparation-tools > .semi-select {
min-width: 130px;
}
.preparation-table {
min-height: 0;
width: 100%;
}
.preparation-tools .semi-button {
white-space: nowrap;
}
.preparation-fields {
display: flex;
flex-direction: column;
gap: 8px;
min-height: 0;
}
.preparation-filters {
display: grid;
grid-template-columns: repeat(auto-fit, minmax(160px, 1fr));
gap: 10px;
margin: 10px 0;
}
.preparation-filters label {
display: flex;
flex-direction: column;
gap: 4px;
font-size: 12px;
color: var(--semi-color-text-2);
}
.preparation-detail {
padding: 4px;
display: flex;
flex-direction: column;
gap: 12px;
}
.preparation-detail > label {
display: flex;
flex-direction: column;
gap: 6px;
}
@media (max-width: 700px) {
.preparation-page {
padding: 8px;
}
.preparation-tools > .semi-input-wrapper {
width: 100%;
}
}
+62
View File
@@ -0,0 +1,62 @@
export type Scope = {
instrument_type: string;
region: string;
universe: string;
delay: number;
};
export const defaultScope: Scope = {
instrument_type: "EQUITY",
region: "USA",
universe: "TOP3000",
delay: 1,
};
export const scopeKey = (s: Scope) =>
`${s.instrument_type}|${s.region}|${s.universe}|${s.delay}`;
export const scopeLabel = (s: Scope) =>
`${s.region} / ${s.universe} / Delay ${s.delay}`;
export type FieldRecord = {
id: string;
field_id: string;
name: string | null;
description: string | null;
dataset_id: string;
dataset_name: string;
field_type: string | null;
scope: Scope;
source: "local" | "worldquant";
collection_version: string | null;
coverage: number | null;
user_count: number | null;
alpha_count: number | null;
category?: string | null;
subcategory?: string | null;
synced_at: string | null;
fetched_at: string;
};
export type Preparation = {
id: string;
name: string;
note: string;
scope: Scope;
version: number;
field_count: number;
dataset_count: number;
created_at: string;
updated_at: string;
};
export type PreparationRef = { id: string; version: number };
export type Page<T> = {
items: T[];
total: number;
limit: number;
offset: number;
has_more?: boolean;
total_known?: boolean;
};
export const fieldReference = (f: FieldRecord) => ({
field_id: f.id,
dataset_id: f.dataset_id,
scope: f.scope,
source: f.source,
collection_version: f.collection_version,
});
+1 -1
View File
@@ -222,7 +222,7 @@ export function ExperimentView({
}) })
} }
> >
{input.dataset_id} · {input.scope.region}/{input.scope.universe}/D {input.name} · {input.scope.region}/{input.scope.universe}/D
{input.scope.delay} {input.scope.delay}
</Button> </Button>
))} ))}
+17 -17
View File
@@ -1,3 +1,7 @@
import {
ResearchDataInput,
researchSelection,
} from "../preparations/ResearchDataInput";
import { ResearchList, ResearchDetail } from "./ResearchList"; import { ResearchList, ResearchDetail } from "./ResearchList";
import { useEffect, useRef, useState } from "react"; import { useEffect, useRef, useState } from "react";
import { import {
@@ -66,9 +70,7 @@ export function FeaturesPage({
`/research/assets?kind=feature&q=${encodeURIComponent(q)}&offset=${(page - 1) * 25}`, `/research/assets?kind=feature&q=${encodeURIComponent(q)}&offset=${(page - 1) * 25}`,
{ signal: controller.signal }, { signal: controller.signal },
), ),
api<{ items: InputSnapshot[] }>("/research/inputs", { Promise.resolve({ items: inputs }),
signal: controller.signal,
}),
]) ])
.then(([list, fixed]) => { .then(([list, fixed]) => {
setItems(list.items); setItems(list.items);
@@ -128,7 +130,7 @@ export function FeaturesPage({
body: JSON.stringify({ body: JSON.stringify({
kind: "feature", kind: "feature",
version: asset?.version, version: asset?.version,
content: draft, content: { ...draft, ...researchSelection(draft.input_ids, inputs) },
}), }),
}, },
); );
@@ -317,7 +319,7 @@ export function FeaturesPage({
.filter((i) => draft.input_ids.includes(i.id)) .filter((i) => draft.input_ids.includes(i.id))
.map((i) => ({ .map((i) => ({
id: i.id, id: i.id,
name: `${i.dataset_id} · ${i.scope.region}/${i.scope.universe}/D${i.scope.delay} · ${i.field_ids.length} 字段 · ${i.id}`, name: `${i.name} · ${i.scope.region}/${i.scope.universe}/D${i.scope.delay} · ${i.field_ids.length} 字段 · ${i.id}`,
}))} }))}
disabled={!!busy} disabled={!!busy}
busy={busy === "generate"} busy={busy === "generate"}
@@ -330,7 +332,7 @@ export function FeaturesPage({
await post<FeatureAsset>("/research/generate", { await post<FeatureAsset>("/research/generate", {
name: draft.name, name: draft.name,
hypothesis: draft.hypothesis, hypothesis: draft.hypothesis,
input_ids: draft.input_ids, ...researchSelection(draft.input_ids, inputs),
method: "feature", method: "feature",
}), }),
); );
@@ -347,23 +349,21 @@ export function FeaturesPage({
</section> </section>
<label> <label>
固定输入 固定输入
<ResearchSelect <ResearchDataInput
label="特征固定输入" label="选择特征数据准备"
multiple ids={draft.input_ids}
filter inputs={inputs}
value={draft.input_ids} onChange={(ids, rows) => {
optionList={inputs.map((i) => ({ setInputs(rows);
value: i.id, setDraft({ ...draft, input_ids: ids });
label: `${i.dataset_id} · ${i.scope.region}/${i.scope.universe}/D${i.scope.delay} · ${i.field_ids.length} 字段 · ${i.id.slice(0, 8)}`, }}
}))}
onChange={(v) => setDraft({ ...draft, input_ids: v as string[] })}
/> />
</label> </label>
{inputs {inputs
.filter((i) => draft.input_ids.includes(i.id)) .filter((i) => draft.input_ids.includes(i.id))
.map((i) => ( .map((i) => (
<details key={i.id}> <details key={i.id}>
<summary>{i.dataset_id} 的固定字段</summary> <summary>{i.name} 的固定字段</summary>
<p className="muted">{i.field_ids.join("、")}</p> <p className="muted">{i.field_ids.join("、")}</p>
</details> </details>
))} ))}
+24 -21
View File
@@ -1,3 +1,7 @@
import {
ResearchDataInput,
researchSelection,
} from "../preparations/ResearchDataInput";
import { useEffect, useState } from "react"; import { useEffect, useState } from "react";
import { import {
Banner, Banner,
@@ -54,7 +58,7 @@ export function FlowLaunchForm({
useEffect(() => { useEffect(() => {
const c = new AbortController(); const c = new AbortController();
Promise.all([ Promise.all([
api<{ items: InputSnapshot[] }>("/research/inputs", { signal: c.signal }), Promise.resolve({ items: inputs }),
api<{ items: Asset[] }>("/research/assets?kind=template&limit=100", { api<{ items: Asset[] }>("/research/assets?kind=template&limit=100", {
signal: c.signal, signal: c.signal,
}), }),
@@ -73,7 +77,7 @@ export function FlowLaunchForm({
setConfirmation({ setConfirmation({
request_id: crypto.randomUUID(), request_id: crypto.randomUUID(),
name, name,
input_ids: ids, ...researchSelection(ids, inputs),
hypothesis, hypothesis,
settings, settings,
budget, budget,
@@ -133,22 +137,17 @@ export function FlowLaunchForm({
</label> </label>
<label> <label>
固定数据范围 固定数据范围
<ResearchSelect <ResearchDataInput
label="自动研究固定输入" label="选择研究数据准备"
multiple ids={ids}
filter inputs={inputs}
value={ids} onChange={(ids, rows) => {
optionList={inputs.map((i) => ({ setInputs(rows);
value: i.id, setIds(ids);
label: `${i.dataset_id} · ${i.scope.region}/${i.scope.universe}/D${i.scope.delay} · ${i.id.slice(0, 8)}`, const first = rows.find((r) => r.id === ids[0]);
}))}
onChange={(v) => {
const next = v as string[];
setIds(next);
const first = inputs.find((i) => i.id === next[0]);
if (first) if (first)
setSettings((s) => ({ setSettings((old) => ({
...s, ...old,
region: first.scope.region, region: first.scope.region,
universe: first.scope.universe, universe: first.scope.universe,
delay: first.scope.delay, delay: first.scope.delay,
@@ -290,8 +289,10 @@ export function FlowLaunchForm({
</p> </p>
)} )}
<p> <p>
{confirmation.input_ids.length} 个固定输入 · {settings.region}/ {confirmation.input_ids.length +
{settings.universe}/D{settings.delay} (confirmation.preparation_refs?.length ?? 0)}{" "}
个固定输入 · {settings.region}/{settings.universe}/D
{settings.delay}
</p> </p>
<ul> <ul>
{confirmation.input_ids.map((id) => { {confirmation.input_ids.map((id) => {
@@ -299,7 +300,7 @@ export function FlowLaunchForm({
return ( return (
<li key={id}> <li key={id}>
{input {input
? `${input.dataset_id} · ${input.scope.region}/${input.scope.universe}/D${input.scope.delay}` ? `${input.name} · ${input.scope.region}/${input.scope.universe}/D${input.scope.delay}`
: id}{" "} : id}{" "}
· {id.slice(0, 8)} · {id.slice(0, 8)}
</li> </li>
@@ -315,7 +316,9 @@ export function FlowLaunchForm({
{confirmation.template_id && ( {confirmation.template_id && (
<p> <p>
初始模板: 初始模板:
{templates.find((t) => t.id === confirmation.template_id)?.name}{" "} {
templates.find((t) => t.id === confirmation.template_id)?.name
}{" "}
· v{confirmation.template_version} · v{confirmation.template_version}
</p> </p>
)} )}
+31 -28
View File
@@ -1,3 +1,7 @@
import {
ResearchDataInput,
researchSelection,
} from "../preparations/ResearchDataInput";
import { ResearchList, ResearchDetail } from "./ResearchList"; import { ResearchList, ResearchDetail } from "./ResearchList";
import { ResearchSelect } from "./ResearchSelect"; import { ResearchSelect } from "./ResearchSelect";
import { useEffect, useRef, useState } from "react"; import { useEffect, useRef, useState } from "react";
@@ -101,7 +105,7 @@ export function ResearchWorkspace({
api<{ items: Asset[]; total: number }>( api<{ items: Asset[]; total: number }>(
`/research/assets?kind=template&limit=25&offset=${assetPage * 25}&q=${encodeURIComponent(search)}`, `/research/assets?kind=template&limit=25&offset=${assetPage * 25}&q=${encodeURIComponent(search)}`,
), ),
api<{ items: InputSnapshot[] }>("/research/inputs"), Promise.resolve({ items: inputs }),
api<{ items: typeof history; total: number }>( api<{ items: typeof history; total: number }>(
`/research/experiments?limit=25&offset=${historyPage * 25}`, `/research/experiments?limit=25&offset=${historyPage * 25}`,
), ),
@@ -129,7 +133,11 @@ export function ResearchWorkspace({
research_asset_id: asset?.id, research_asset_id: asset?.id,
research_experiment_id: experiment?.id, research_experiment_id: experiment?.id,
alpha_id: parent.split(/[,,\s]+/)[0] || undefined, alpha_id: parent.split(/[,,\s]+/)[0] || undefined,
template_input_id: inputIds.length === 1 ? inputIds[0] : undefined, input_snapshot_id:
inputIds.length === 1 &&
!inputs.find((i) => i.id === inputIds[0])?.preparation_ref
? inputIds[0]
: undefined,
}); });
}, [active, page, asset?.id, experiment?.id, parent, inputIds, onContext]); }, [active, page, asset?.id, experiment?.id, parent, inputIds, onContext]);
useEffect(() => { useEffect(() => {
@@ -180,17 +188,6 @@ export function ResearchWorkspace({
setMethod("structure"); setMethod("structure");
} }
}, [active, action]); }, [active, action]);
function selectInputs(ids: string[]) {
setInputIds(ids);
const first = inputs.find((item) => item.id === ids[0]);
if (first)
setSettings((old) => ({
...old,
region: first.scope.region,
universe: first.scope.universe,
delay: first.scope.delay,
}));
}
async function refreshSettings() { async function refreshSettings() {
const data = await post<{ const data = await post<{
content: { content: {
@@ -236,7 +233,7 @@ export function ResearchWorkspace({
const next = await post<Asset>("/research/generate", { const next = await post<Asset>("/research/generate", {
name: template.name, name: template.name,
hypothesis, hypothesis,
input_ids: inputIds, ...researchSelection(inputIds, inputs),
parent_alpha_ids: parent ? parent.split(/[,,\s]+/).filter(Boolean) : [], parent_alpha_ids: parent ? parent.split(/[,,\s]+/).filter(Boolean) : [],
method: page === "variants" ? "structure" : "template", method: page === "variants" ? "structure" : "template",
}); });
@@ -274,7 +271,7 @@ export function ResearchWorkspace({
const next = await post<Experiment>("/research/experiments", { const next = await post<Experiment>("/research/experiments", {
asset_id: asset!.id, asset_id: asset!.id,
version: asset!.version, version: asset!.version,
input_ids: inputIds, ...researchSelection(inputIds, inputs),
hypothesis, hypothesis,
settings, settings,
mode, mode,
@@ -580,20 +577,26 @@ export function ResearchWorkspace({
)} )}
<label> <label>
固定研究输入 固定研究输入
<ResearchSelect <ResearchDataInput
multiple label="选择数据准备"
filter ids={inputIds}
label="固定研究输入" inputs={inputs}
value={inputIds} onChange={(ids, rows) => {
optionList={inputs.map((input) => ({ setInputs(rows);
value: input.id, setInputIds(ids);
label: `${input.dataset_id} · ${input.scope.region}/${input.scope.universe}/D${input.scope.delay} · ${input.field_ids.length} 字段 · ${input.id.slice(0, 8)}`, const first = rows.find((r) => r.id === ids[0]);
}))} if (first)
onChange={(value) => selectInputs(value as string[])} setSettings((old) => ({
...old,
region: first.scope.region,
universe: first.scope.universe,
delay: first.scope.delay,
}));
}}
/> />
</label> </label>
<p className="research-hint"> <p className="research-hint">
在数据目录中保存字段选择。跨数据集分别关联输入,跨市场使用目标范围的独立输入。 从数据准备选择集合;每个集合可包含同范围下多个数据集的字段。
</p> </p>
{page === "variants" && method === "settings" ? ( {page === "variants" && method === "settings" ? (
<label> <label>
@@ -624,7 +627,7 @@ export function ResearchWorkspace({
references={[ references={[
...selectedInputs.map((input) => ({ ...selectedInputs.map((input) => ({
id: input.id, id: input.id,
name: `${input.dataset_id} · ${input.scope.region}/${input.scope.universe}/D${input.scope.delay} · ${input.field_ids.length} 字段 · ${input.id}`, name: `${input.name} · ${input.scope.region}/${input.scope.universe}/D${input.scope.delay} · ${input.field_ids.length} 字段 · ${input.id}`,
})), })),
...(page === "variants" && parent ...(page === "variants" && parent
? [ ? [
@@ -857,7 +860,7 @@ export function ResearchWorkspace({
setExperiment( setExperiment(
await post("/research/variants/settings", { await post("/research/variants/settings", {
alpha_id: parent, alpha_id: parent,
input_ids: inputIds, ...researchSelection(inputIds, inputs),
...(hypothesis ? { hypothesis } : {}), ...(hypothesis ? { hypothesis } : {}),
}), }),
); );
+12 -3
View File
@@ -62,19 +62,28 @@ export function SourceDetails({
打开研究会话 打开研究会话
</Button> </Button>
)} )}
{source.template_input_id && ( {[
...new Set([
...(source.input_snapshot_ids ?? []),
...(source.input_snapshot_id ? [source.input_snapshot_id] : []),
]),
].map((id, index) => (
<Button <Button
key={id}
onClick={() => onClick={() =>
onAction({ onAction({
type: "open_research_input", type: "open_research_input",
input_id: source.template_input_id!, input_id: id,
nonce: Date.now(), nonce: Date.now(),
}) })
} }
> >
查看研究输入 查看研究输入
{(source.input_snapshot_ids?.length ?? 0) > 1
? ` ${index + 1}`
: ""}
</Button> </Button>
)} ))}
{source.parent_run_id && ( {source.parent_run_id && (
<Button <Button
onClick={() => onClick={() =>
+1
View File
@@ -21,6 +21,7 @@ export type FlowLaunch = {
request_id: string; request_id: string;
name: string; name: string;
input_ids: string[]; input_ids: string[];
preparation_refs?: { id: string; version: number }[];
hypothesis: string; hypothesis: string;
settings: SimulationSettings; settings: SimulationSettings;
budget: Budget; budget: Budget;
+7 -2
View File
@@ -1,3 +1,4 @@
import type { FieldRecord, PreparationRef } from "../preparations/types";
import type { Candidate } from "../backtests/types"; import type { Candidate } from "../backtests/types";
export type Variable = { export type Variable = {
kind: kind:
@@ -43,10 +44,14 @@ export type InputSnapshot = {
universe: string; universe: string;
delay: 0 | 1; delay: 0 | 1;
}; };
dataset_id: string; name: string;
dataset_ids: string[];
fields: FieldRecord[];
preparation_ref?: PreparationRef;
field_ids: string[]; field_ids: string[];
field_types: Record<string, string>; field_types: Record<string, string>;
collection_version: string; preparation_id?: string;
preparation_version?: number;
created_at: string; created_at: string;
}; };
export type ResearchCandidate = Candidate & { export type ResearchCandidate = Candidate & {
+3
View File
@@ -139,6 +139,9 @@ export type Job = {
}; };
checkpoint: { checkpoint: {
date?: string; date?: string;
dataset_id?: string;
datasets_completed?: number;
datasets_total?: number;
dates_completed?: number; dates_completed?: number;
dates_total?: number; dates_total?: number;
alpha_id?: string; alpha_id?: string;
+105 -248
View File
@@ -1,14 +1,12 @@
import { test, expect, type Page } from "@playwright/test"; import { test, expect, type Page } from "@playwright/test";
import { choosePreparation } from "./preparation-helpers";
const headers = { "X-WQ-Request": "1" };
const scope = { const scope = {
instrument_type: "EQUITY", instrument_type: "EQUITY",
region: "USA", region: "USA",
universe: "TOP3000", universe: "TOP3000",
delay: 1, delay: 1,
}; };
const query = new URLSearchParams(
Object.entries(scope).map(([k, v]) => [k, String(v)]),
).toString();
const headers = { "X-WQ-Request": "1" };
async function setup(page: Page) { async function setup(page: Page) {
await page.goto("/#datasets"); await page.goto("/#datasets");
await page.getByLabel("密码", { exact: true }).fill("browser-test-password"); await page.getByLabel("密码", { exact: true }).fill("browser-test-password");
@@ -20,273 +18,132 @@ async function setup(page: Page) {
headers, headers,
data: { email: "test@example.com", password: "synthetic-only" }, data: { email: "test@example.com", password: "synthetic-only" },
}); });
const connect = await ( const job = await (
await page.request.post("/api/v1/account/connect", { headers }) await page.request.post("/api/v1/account/connect", { headers })
).json(); ).json();
await expect await expect
.poll( .poll(
async () => async () =>
( (await (await page.request.get(`/api/v1/sync-jobs/${job.id}`)).json())
await ( .status,
await page.request.get(`/api/v1/sync-jobs/${connect.id}`)
).json()
).status,
) )
.toBe("completed"); .toBe("completed");
await page.getByRole("button", { name: "同步目录", exact: true }).click(); await page.reload();
}
async function sync(page: Page, dataset_id: string | null) {
const job = await (
await page.request.post("/api/v1/catalog/sync-jobs", {
headers,
data: { scope, dataset_id },
})
).json();
await expect
.poll(
async () =>
(await (await page.request.get(`/api/v1/sync-jobs/${job.id}`)).json())
.status,
)
.toBe("completed");
}
test("数据集完整同步、使用、集合编辑与研究选择", async ({ page }) => {
const errors: string[] = [];
page.on("pageerror", (e) => errors.push(e.message));
await setup(page);
await sync(page, null);
await page.reload();
const row = page.getByRole("row").filter({ hasText: "TEST 财务报表" });
await expect(
row.getByRole("button", { name: "使用", exact: true }),
).toBeDisabled();
await row.getByRole("button", { name: "同步", exact: true }).click();
await expect await expect
.poll( .poll(
async () => async () =>
( (
await ( await (
await page.request.get(`/api/v1/catalog/datasets?${query}`) await page.request.get(
"/api/v1/catalog/datasets/TEST_FIN?region=USA&universe=TOP3000&delay=1",
)
).json() ).json()
).total, ).complete_count,
) )
.toBe(3); .toBe(123);
await page.keyboard.press("Escape"); await page.keyboard.press("Escape");
await expect(
page.getByRole("button", { name: "TEST 财务报表", exact: true }),
).toBeVisible();
}
async function openFields(page: Page) {
await page
.getByRole("row")
.filter({ hasText: "TEST 财务报表" })
.getByRole("button", { name: "查看字段", exact: true })
.click();
const sync = page
.getByRole("dialog", { name: "数据字段", exact: true })
.getByRole("button", { name: "同步全部字段", exact: true });
if (await sync.isVisible()) {
await sync.click();
await expect
.poll(
async () =>
(
await (
await page.request.get(
`/api/v1/catalog/datasets/TEST_FIN/fields?${query}`,
)
).json()
).complete_count,
)
.toBe(123);
await page.keyboard.press("Escape");
}
}
test("目录筛选、双层详情、完整输入、排除、备注与刷新", async ({ page }) => {
const errors: string[] = [];
page.on("pageerror", (e) => errors.push(e.message));
await page.setViewportSize({ width: 1440, height: 1000 });
await setup(page);
await page.getByLabel("分类", { exact: true }).click();
await page.getByRole("option", { name: /基本面/ }).click();
await page.getByLabel("子分类", { exact: true }).click();
await page.getByRole("option", { name: /财务报表/ }).click();
await openFields(page);
const fields = page.getByRole("dialog", { name: "数据字段", exact: true });
await expect(fields).toBeVisible();
await expect(fields).toContainText("全部 123 个字段");
expect((await fields.boundingBox())!.width).toBeCloseTo(864, 0);
await page.getByLabel("搜索字段").fill("字段 12");
await expect(fields.getByRole("button", { name: /TEST 字段/ })).toHaveCount(
3,
);
await page
.getByRole("button", { name: "TEST 字段 122", exact: true })
.click();
const detail = page.getByRole("dialog", { name: "字段详情", exact: true });
await expect(detail).toContainText("FUTURE_TYPE");
expect((await detail.boundingBox())!.width).toBeCloseTo(864, 0);
await page.getByLabel("研究备注", { exact: true }).fill("保留字段研究假设");
await page.getByRole("button", { name: "保存备注", exact: true }).click();
await expect(page.getByText("研究备注已保存")).toBeVisible();
await page.keyboard.press("Escape");
await expect(detail).not.toBeVisible();
await expect(page.getByLabel("搜索字段")).toHaveValue("字段 12");
await expect(
page.getByRole("button", { name: "TEST 字段 122", exact: true }),
).toBeFocused();
await page
.getByRole("button", { name: "用于 Alpha 模板", exact: true })
.click();
await page.getByRole("button", { name: "保存输入草稿", exact: true }).click();
await expect(
page.getByRole("heading", { name: "输入草稿已保存" }),
).toBeVisible();
expect(
(
await (await page.request.get(`/api/v1/catalog/inputs?${query}`)).json()
)[0].field_ids,
).toHaveLength(123);
await page.keyboard.press("Escape");
await page
.getByRole("checkbox", { name: "选择TEST 字段 122", exact: true })
.press("Space");
await page.getByRole("button", { name: "覆盖率排序" }).click();
await expect(fields).toContainText("122 / 123 个字段");
await page
.getByRole("button", { name: "用于 Alpha 模板", exact: true })
.click();
await page.getByRole("button", { name: "保存输入草稿", exact: true }).click();
await expect(
page.getByRole("heading", { name: "输入草稿已保存" }),
).toBeVisible();
const subset = (
await (await page.request.get(`/api/v1/catalog/inputs?${query}`)).json()
)[0];
expect(subset.field_ids).toHaveLength(122);
expect(subset.field_ids).not.toContain("TEST_FIN_122");
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "恢复全选" }).click();
await page
.getByRole("checkbox", { name: "选择本数据集全部字段" })
.press("Space");
await expect(
page.getByRole("button", { name: "用于 Alpha 模板", exact: true }),
).toBeDisabled();
await page.getByRole("button", { name: "恢复全选" }).click();
await page
.getByRole("button", { name: "TEST 字段 122", exact: true })
.click();
await expect(page.getByLabel("研究备注", { exact: true })).toHaveValue(
"保留字段研究假设",
);
await page.screenshot({ path: "../output/playwright/dataset-desktop.png" });
await page.keyboard.press("Escape");
await page.keyboard.press("Escape");
await expect(page.getByLabel("分类", { exact: true })).toContainText(
"基本面",
);
await page.reload(); await page.reload();
await row.getByRole("button", { name: "查看", exact: true }).click();
const detail = page.getByRole("dialog", {
name: "TEST 财务报表",
exact: true,
});
await expect(detail).toContainText("TEST_FIN_001");
await detail.getByLabel("研究备注").fill("验证备注");
await detail.getByRole("button", { name: "保存备注", exact: true }).click();
await page.keyboard.press("Escape");
await row.getByRole("button", { name: "使用", exact: true }).click();
await page.getByRole("link", { name: "打开数据准备" }).click();
const prep = page.getByRole("dialog", { name: "数据准备详情", exact: true });
await expect(prep).toContainText("123 字段 / 1 数据集");
await prep.getByLabel("集合名称", { exact: true }).fill("浏览器准备集合");
await prep.getByRole("button", { name: "保存名称和备注" }).click();
await expect(prep).toContainText("v2");
await prep
.getByRole("row")
.filter({ hasText: "TEST_FIN_000" })
.locator("span.semi-checkbox")
.click();
await prep.getByRole("button", { name: "移除所选字段" }).click();
await expect(prep).toContainText("122 字段");
await page.keyboard.press("Escape");
await page
.getByRole("navigation", { name: "主导航" })
.getByRole("button", { name: "特征工程", exact: true })
.click();
await page.getByRole("button", { name: "新建特征方案", exact: true }).click();
await choosePreparation(page, "浏览器准备集合", "选择特征数据准备");
await expect( await expect(
page.getByRole("button", { name: "已保存输入 (2)" }), page.locator("p").filter({ hasText: /浏览器准备集合 ·.*122 字段/ }),
).toBeVisible(); ).toBeVisible();
expect(errors).toEqual([]); expect(errors).toEqual([]);
}); });
test("在线字段跨页多选直接准备,本地只展示完整同步字段", async ({ page }) => {
test("窄屏抽屉、键盘隔离和研究范围切换", async ({ page }) => {
await page.setViewportSize({ width: 390, height: 844 });
await setup(page); await setup(page);
await openFields(page);
const fields = page.getByRole("dialog", { name: "数据字段", exact: true });
await expect(fields).toBeVisible();
expect((await fields.boundingBox())!.width).toBeCloseTo(234, 0);
await page await page
.getByRole("button", { name: "TEST 字段 000", exact: true }) .getByRole("navigation", { name: "主导航" })
.getByRole("button", { name: "字段目录", exact: true })
.click(); .click();
await expect( const online = page.getByRole("tabpanel", { name: "worldquant接口" });
page.getByRole("dialog", { name: "字段详情", exact: true }), await expect(online.getByText("TEST_FIN_001", { exact: true })).toBeVisible();
).toBeVisible(); await online
await page.keyboard.press("Escape"); .getByRole("row")
await expect(fields).toBeVisible(); .filter({ hasText: "TEST_FIN_001" })
expect( .locator("span.semi-checkbox")
await page.evaluate(
() => document.documentElement.scrollWidth <= innerWidth,
),
).toBe(true);
await page.screenshot({ path: "../output/playwright/dataset-mobile.png" });
await page.keyboard.press("Escape");
await page.getByLabel("Region", { exact: true }).click();
await page.getByRole("option", { name: /CHN/ }).click();
await expect(page.getByLabel("Universe", { exact: true })).toContainText(
"TOP2000",
);
await expect(
page.getByRole("button", { name: "用于 Alpha 模板", exact: true }),
).toBeDisabled();
});
test("非首页排除、聊天暂存详情、遮罩逐层关闭和多尺寸", async ({ page }) => {
const errors: string[] = [];
page.on("pageerror", (e) => errors.push(e.message));
await page.setViewportSize({ width: 1280, height: 900 });
await setup(page);
await openFields(page);
const fields = page.getByRole("dialog", { name: "数据字段", exact: true });
await fields.getByRole("button", { name: "Next", exact: true }).click();
await expect(
page.getByRole("button", { name: "TEST 字段 025", exact: true }),
).toBeVisible();
await page
.getByRole("checkbox", { name: "选择TEST 字段 025", exact: true })
.press("Space");
await page.getByLabel("搜索字段").fill("字段 12");
await expect(fields).toContainText("122 / 123 个字段");
await page.getByLabel("字段类型", { exact: true }).click();
await page.getByRole("option", { name: /VECTOR/ }).click();
await expect(
page.getByRole("button", { name: "TEST 字段 120", exact: true }),
).toBeVisible();
await page
.getByRole("button", { name: "TEST 字段 120", exact: true })
.click(); .click();
await page.getByLabel("研究备注", { exact: true }).fill("未保存草稿需要恢复"); await online.locator(".semi-page-item").filter({ hasText: /^2$/ }).click();
await page await online
.getByRole("dialog", { name: "字段详情", exact: true }) .getByRole("row")
.getByRole("button", { name: "AI 助手", exact: true }) .filter({ hasText: "TEST_FIN_025" })
.locator("span.semi-checkbox")
.click(); .click();
await expect(fields).not.toBeVisible(); await expect(online.getByText(/已选 2 个字段/)).toBeVisible();
await page.keyboard.press("Escape"); await online
await expect(page.getByLabel("研究备注", { exact: true })).toHaveValue( .getByRole("button", { name: "加入数据准备", exact: true })
"未保存草稿需要恢复", .click();
); const add = page.getByRole("dialog", { name: "加入数据准备", exact: true });
for (const width of [850, 390, 1920]) { await add.getByLabel("新集合名称").fill("在线选取集合");
await page.setViewportSize({ width, height: 900 }); await add.getByRole("button", { name: "新建集合并加入" }).click();
await expect( await expect(add).not.toBeVisible();
page.getByRole("button", { name: "关闭详情", exact: true }), const collections = await (
).toBeVisible(); await page.request.get("/api/v1/data-preparations?q=在线选取集合")
expect( ).json();
await page.evaluate( expect(collections.items[0].field_count).toBe(2);
() => document.documentElement.scrollWidth <= innerWidth, await page.getByRole("tab", { name: "本地同步", exact: true }).click();
), const local = page.getByRole("tabpanel", { name: "本地同步" });
).toBe(true); // Other browser tests share this isolated server, so verify the unsynced dataset explicitly.
} await local.getByLabel("字段数据集").fill("TEST_NEWS");
await page.setViewportSize({ width: 1440, height: 900 }); await local.getByRole("button", { name: "查询", exact: true }).click();
await page.mouse.click(10, 400); await expect(local.getByText("共 0 个字段", { exact: true })).toBeVisible();
await expect( await page.screenshot({
page.getByRole("dialog", { name: "字段详情", exact: true }), path: "test-results/preparations-fields.png",
).not.toBeVisible(); fullPage: true,
await expect(fields).toBeVisible(); });
await expect(fields).toContainText("122 / 123 个字段");
await page.mouse.click(10, 400);
await expect(fields).not.toBeVisible();
expect(errors).toEqual([]);
});
test("平台选项失败可重试,地区联动限制 Delay 和 Universe", async ({ page }) => {
await setup(page);
await page.getByLabel("Delay", { exact: true }).click();
await page.getByRole("option", { name: /Delay 0/ }).click();
await page.getByLabel("Region", { exact: true }).click();
await page.getByRole("option", { name: /IND/ }).click();
await expect(page.getByLabel("Universe", { exact: true })).toContainText(
"TOP500",
);
await expect(page.getByLabel("Delay", { exact: true })).toContainText(
"Delay 1",
);
await page.getByLabel("Delay", { exact: true }).click();
await expect(page.getByRole("option", { name: /Delay 0/ })).toHaveCount(0);
await page.keyboard.press("Escape");
await page.route("**/api/v1/catalog/scopes", (route) =>
route.fulfill({
status: 502,
contentType: "application/json",
body: JSON.stringify({ detail: "测试平台选项不可用" }),
}),
);
await page.getByRole("button", { name: "Alpha 管理", exact: true }).click();
await page.getByRole("button", { name: "数据目录", exact: true }).click();
await expect(page.getByText("测试平台选项不可用")).toBeVisible();
await expect(
page.getByRole("button", { name: "同步目录", exact: true }),
).toBeDisabled();
await page.unroute("**/api/v1/catalog/scopes");
await page.getByRole("button", { name: "重试获取选项" }).click();
await expect(
page.getByRole("button", { name: "同步目录", exact: true }),
).toBeEnabled();
}); });
+21
View File
@@ -0,0 +1,21 @@
import { expect, type Page } from "@playwright/test";
export async function choosePreparation(
page: Page,
name: string,
label = "选择研究数据准备",
) {
await page.getByRole("button", { name: label, exact: true }).click();
const picker = page.getByRole("dialog", {
name: "选择数据准备集合",
exact: true,
});
await picker.getByLabel("搜索数据准备集合").fill(name);
await picker
.getByRole("row")
.filter({ hasText: name })
.first()
.locator("span.semi-checkbox")
.click();
await picker.getByRole("button", { name: "选择", exact: true }).click();
await expect(picker).not.toBeVisible();
}
+3 -7
View File
@@ -1,3 +1,4 @@
import { choosePreparation } from "./preparation-helpers";
import { expect, test } from "@playwright/test"; import { expect, test } from "@playwright/test";
const headers = { "X-WQ-Request": "1" }; const headers = { "X-WQ-Request": "1" };
const scope = { const scope = {
@@ -71,13 +72,12 @@ test("native QuantFlow canvas saves versions, branches and restores artifacts",
await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`) await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`)
).json(); ).json();
const fixed = await ( const fixed = await (
await page.request.post("/api/v1/catalog/inputs", { await page.request.post("/api/v1/data-preparations/from-dataset", {
headers, headers,
data: { data: {
scope, scope,
dataset_id: "TEST_FIN", dataset_id: "TEST_FIN",
collection_version: fields.collection_version, collection_version: fields.collection_version,
selection: "all",
}, },
}) })
).json(); ).json();
@@ -175,11 +175,7 @@ test("native QuantFlow canvas saves versions, branches and restores artifacts",
.getByRole("button", { name: "启动已保存流程", exact: true }) .getByRole("button", { name: "启动已保存流程", exact: true })
.click(); .click();
await page.getByLabel("自动研究假设").fill("画布条件分支与固定版本"); await page.getByLabel("自动研究假设").fill("画布条件分支与固定版本");
await page.getByRole("combobox", { name: "自动研究固定输入" }).click(); await choosePreparation(page, fixed.name, "选择研究数据准备");
await page
.getByRole("option")
.filter({ hasText: fixed.id.slice(0, 8) })
.click();
await page.getByLabel("自动研究假设").click(); await page.getByLabel("自动研究假设").click();
for (const [label, value] of [ for (const [label, value] of [
["最大轮数", "1"], ["最大轮数", "1"],
+3 -7
View File
@@ -1,3 +1,4 @@
import { choosePreparation } from "./preparation-helpers";
import { expect, test } from "@playwright/test"; import { expect, test } from "@playwright/test";
const headers = { "X-WQ-Request": "1" }; const headers = { "X-WQ-Request": "1" };
const scope = { const scope = {
@@ -71,13 +72,12 @@ test("finite pipeline confirms scope and budget, executes native steps and resto
await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`) await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`)
).json(); ).json();
const fixed = await ( const fixed = await (
await page.request.post("/api/v1/catalog/inputs", { await page.request.post("/api/v1/data-preparations/from-dataset", {
headers, headers,
data: { data: {
scope, scope,
dataset_id: "TEST_FIN", dataset_id: "TEST_FIN",
collection_version: fields.collection_version, collection_version: fields.collection_version,
selection: "all",
}, },
}) })
).json(); ).json();
@@ -97,11 +97,7 @@ test("finite pipeline confirms scope and budget, executes native steps and resto
await page.getByRole("button", { name: "新建自动研究", exact: true }).click(); await page.getByRole("button", { name: "新建自动研究", exact: true }).click();
await page.getByLabel("自动研究名称").fill("浏览器有限研究"); await page.getByLabel("自动研究名称").fill("浏览器有限研究");
await page.getByLabel("自动研究假设").fill("研究排名稳定性"); await page.getByLabel("自动研究假设").fill("研究排名稳定性");
await page.getByRole("combobox", { name: "自动研究固定输入" }).click(); await choosePreparation(page, fixed.name, "选择研究数据准备");
await page
.getByRole("option")
.filter({ hasText: fixed.id.slice(0, 8) })
.click();
await page.getByLabel("自动研究名称").click(); await page.getByLabel("自动研究名称").click();
for (const [label, value] of [ for (const [label, value] of [
["最大轮数", "1"], ["最大轮数", "1"],
+23 -20
View File
@@ -68,17 +68,27 @@ test("fixed input, AI approval, backtest and research source remain connected",
const fields = await ( const fields = await (
await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`) await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`)
).json(); ).json();
const saved = await page.request.post("/api/v1/catalog/inputs", { const saved = await page.request.post(
headers, "/api/v1/data-preparations/from-dataset",
data: { {
scope, headers,
dataset_id: "TEST_FIN", data: {
collection_version: fields.collection_version, scope,
selection: "all", dataset_id: "TEST_FIN",
collection_version: fields.collection_version,
},
}, },
}); );
expect(saved.status()).toBe(201); expect(saved.status()).toBe(201);
const input = await saved.json(); const preparation = await saved.json();
const input = (
await (
await page.request.post("/api/v1/data-preparations/freeze", {
headers,
data: { items: [{ id: preparation.id, version: preparation.version }] },
})
).json()
).items[0];
const conversation = await ( const conversation = await (
await page.request.post("/api/v1/ai/conversations", { headers }) await page.request.post("/api/v1/ai/conversations", { headers })
).json(); ).json();
@@ -96,7 +106,7 @@ test("fixed input, AI approval, backtest and research source remain connected",
page: "datasets", page: "datasets",
catalog_scope: scope, catalog_scope: scope,
dataset_id: "TEST_FIN", dataset_id: "TEST_FIN",
template_input_id: input.id, input_snapshot_id: input.id,
}, },
}, },
}, },
@@ -130,17 +140,10 @@ test("fixed input, AI approval, backtest and research source remain connected",
.click(); .click();
await expect(chat).not.toBeVisible(); await expect(chat).not.toBeVisible();
const inputSheet = page.getByRole("dialog", { const inputSheet = page.getByRole("dialog", {
name: "模板输入草稿", name: "研究输入快照",
exact: true, exact: true,
}); });
await expect( await expect(inputSheet).toContainText("TEST_FIN_001");
inputSheet.getByRole("heading", { name: "输入草稿已保存", exact: true }), await expect(inputSheet).toContainText("TEST_FIN");
).toBeVisible();
await inputSheet
.getByRole("button", { name: "用此输入研究", exact: true })
.click();
await expect(
chat.locator(".research-ai-reference").filter({ hasText: "固定输入" }),
).toContainText(input.id);
expect(errors).toEqual([]); expect(errors).toEqual([]);
}); });
+9 -10
View File
@@ -1,3 +1,4 @@
import { choosePreparation } from "./preparation-helpers";
import { expect, test } from "@playwright/test"; import { expect, test } from "@playwright/test";
const headers = { "X-WQ-Request": "1" }; const headers = { "X-WQ-Request": "1" };
const scope = { const scope = {
@@ -50,13 +51,12 @@ test("feature conversion, saved views and immutable evaluations", async ({
await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`) await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`)
).json(); ).json();
const fixed = await ( const fixed = await (
await page.request.post("/api/v1/catalog/inputs", { await page.request.post("/api/v1/data-preparations/from-dataset", {
headers, headers,
data: { data: {
scope, scope,
dataset_id: "TEST_FIN", dataset_id: "TEST_FIN",
collection_version: fields.collection_version, collection_version: fields.collection_version,
selection: "all",
}, },
}) })
).json(); ).json();
@@ -67,11 +67,7 @@ test("feature conversion, saved views and immutable evaluations", async ({
await page await page
.getByLabel("特征经济假设", { exact: true }) .getByLabel("特征经济假设", { exact: true })
.fill("排名可降低异常值影响"); .fill("排名可降低异常值影响");
await page.getByRole("combobox", { name: "特征固定输入" }).click(); await choosePreparation(page, fixed.name, "选择特征数据准备");
await page
.getByRole("option")
.filter({ hasText: fixed.id.slice(0, 8) })
.click();
await page.getByLabel("特征方案名称", { exact: true }).click(); await page.getByLabel("特征方案名称", { exact: true }).click();
await page.getByRole("button", { name: "添加处理步骤" }).click(); await page.getByRole("button", { name: "添加处理步骤" }).click();
await page.getByLabel("步骤 1 名称").fill("横截面排名"); await page.getByLabel("步骤 1 名称").fill("横截面排名");
@@ -95,9 +91,12 @@ test("feature conversion, saved views and immutable evaluations", async ({
"/api/v1/research/assets?kind=template&q=特征转换模板", "/api/v1/research/assets?kind=template&q=特征转换模板",
) )
).json(); ).json();
expect(templates.items[0].provenance.feature.content.input_ids).toEqual([ const snapshot = await (
fixed.id, await page.request.get(
]); `/api/v1/research/input-snapshots/${templates.items[0].provenance.feature.content.input_ids[0]}`,
)
).json();
expect(snapshot.preparation_id).toBe(fixed.id);
await page.keyboard.press("Escape"); await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Alpha 管理", exact: true }).click(); await page.getByRole("button", { name: "Alpha 管理", exact: true }).click();
await page.getByRole("button", { name: "保存为新视图" }).click(); await page.getByRole("button", { name: "保存为新视图" }).click();
+30 -15
View File
@@ -1,9 +1,17 @@
import { choosePreparation } from "./preparation-helpers";
import { expect, test } from "@playwright/test"; import { expect, test } from "@playwright/test";
// UI contract fixtures: no external model call or simulation is started by these tests. // UI contract fixtures: no external model call or simulation is started by these tests.
const fixed = { const fixed = {
id: "semi-fixed-input", id: "semi-fixed-input",
dataset_id: "TEST_FIN", name: "Semi 准备集合",
note: "",
version: 1,
field_count: 1,
dataset_count: 1,
dataset_ids: ["TEST_FIN"],
fields: [],
preparation_ref: { id: "semi-fixed-input", version: 1 },
field_ids: ["TEST_FIN_001"], field_ids: ["TEST_FIN_001"],
field_types: { TEST_FIN_001: "MATRIX" }, field_types: { TEST_FIN_001: "MATRIX" },
collection_version: "fixed-v1", collection_version: "fixed-v1",
@@ -20,8 +28,19 @@ for (const mode of ["templates", "variants", "features"] as const) {
}) => { }) => {
const errors: string[] = []; const errors: string[] = [];
page.on("pageerror", (error) => errors.push(error.message)); page.on("pageerror", (error) => errors.push(error.message));
await page.route("**/api/v1/research/inputs", (route) => await page.route("**/api/v1/data-preparations?*", (route) =>
route.fulfill({ json: { items: [fixed] } }), route.fulfill({
json: { items: [fixed], total: 1, limit: 25, offset: 0 },
}),
);
await page.route(
"**/api/v1/data-preparations/semi-fixed-input/selection?*",
(route) => route.fulfill({ json: fixed }),
);
await page.route(
"**/api/v1/research/input-snapshots/semi-fixed-input",
(route) =>
route.fulfill({ json: { ...fixed, preparation_ref: undefined } }),
); );
await page.route( await page.route(
"**/api/v1/research/assets/semi-generated/versions", "**/api/v1/research/assets/semi-generated/versions",
@@ -49,23 +68,18 @@ for (const mode of ["templates", "variants", "features"] as const) {
exact: true, exact: true,
}); });
await expect(generate).toBeDisabled(); await expect(generate).toBeDisabled();
await sheet await choosePreparation(
.getByRole("combobox", { page,
name: feature ? "特征固定输入" : "固定研究输入", fixed.name,
exact: true, feature ? "选择特征数据准备" : "选择数据准备",
}) );
.click();
await page
.getByRole("option")
.filter({ hasText: fixed.id.slice(0, 8) })
.click();
await input.click(); await input.click();
if (mode === "variants") if (mode === "variants")
await sheet.getByLabel("种子 Alpha", { exact: true }).fill("TEST0001"); await sheet.getByLabel("种子 Alpha", { exact: true }).fill("TEST0001");
const hypothesis = "检验 <收益> 与规模\n第二轮保留原始假设"; const hypothesis = "检验 <收益> 与规模\n第二轮保留原始假设";
await input.fill(hypothesis); await input.fill(hypothesis);
await expect( await expect(
sheet.locator(".research-ai-reference").filter({ hasText: fixed.id }), sheet.locator("p").filter({ hasText: /Semi 准备集合 ·/ }),
).toBeVisible(); ).toBeVisible();
let release!: () => void; let release!: () => void;
const pending = new Promise<void>((resolve) => { const pending = new Promise<void>((resolve) => {
@@ -117,7 +131,8 @@ for (const mode of ["templates", "variants", "features"] as const) {
.poll(() => requestBody) .poll(() => requestBody)
.toMatchObject({ .toMatchObject({
hypothesis, hypothesis,
input_ids: [fixed.id], input_ids: [],
preparation_refs: [fixed.preparation_ref],
method: feature method: feature
? "feature" ? "feature"
: mode === "variants" : mode === "variants"