refactor: unify data preparations and research input snapshots
Deploy production / deploy (push) Successful in 53s
Deploy production / deploy (push) Successful in 53s
This commit is contained in:
@@ -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 调度。真实平台过滤/分页协议、范围权限及调度效果需单独联调。迁移按用户确认不兼容旧研究输入,回退需要升级前数据库备份。
|
||||
@@ -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 调度不在本地执行范围。
|
||||
@@ -34,19 +34,23 @@ docker compose ps
|
||||
|
||||
工作空间和 AI 交互统一采用紧凑的 Lark 样式。Alpha 列表只滚动表体,分页保持在可用区域底部;个人信息页独立滚动。
|
||||
|
||||
## 数据集与数据字段
|
||||
## 数据目录、字段目录与数据准备
|
||||
|
||||
从侧栏进入“数据集”,设置 Region、Universe、Delay 后手动同步目录。范围选项表示本版支持的组合,平台账户实际权限以同步结果为准;分类和子分类来自已同步数据。
|
||||
统一流程是“查询或同步字段 → 整理数据准备集合 → 各研究模块选择集合”。
|
||||
|
||||
选中一个数据集后默认使用整集字段;首次使用先同步全部字段。字段列表、搜索、类型、覆盖率、排序及翻页均不改变输入范围,只有明确取消勾选才排除字段。表头选择作用于整个已完成集合,支持恢复全选。字段与详情采用 75% / 30% 的工作区右抽屉,窄屏展开为全宽;逐层关闭保留父层条件。抽屉顶部可打开 AI 助手,业务抽屉暂时隐藏,收起助手后恢复;发送消息时会附带当前范围和输入引用;助手可通过工具读取本地目录和字段,不发送未保存的研究备注。
|
||||
数据目录行操作为“查看”“同步”“使用”。“同步目录”仅刷新数据集清单;行内“同步”更新全部字段。完整分页成功后才发布,失败、取消或刷新期间保留上一完整版本。“使用”复制当前完整字段到新准备集合,首次同步未完成时禁用。
|
||||
|
||||
“用于 Alpha 模板”先保存输入草稿,在服务端固定数据集、研究范围、集合版本、字段 ID 和字段类型。点击“用此输入研究”将该快照带入聊天;也可从“已保存输入”恢复。后续同步不会改变旧草稿。
|
||||
字段目录分为 `worldquant接口` 和 `本地同步`。在线字段可直接加入集合,不标记数据集已同步。本地目录汇总已完整同步字段,提供范围、数据集、类型、关键词、分类、覆盖率、用户数、Alpha 数、同步时间和排序筛选。两种来源支持跨页勾选,表头选择只作用于当前页,切换范围清空选择。
|
||||
|
||||
数据集和字段备注单独保存,版本冲突保留当前草稿。字段同步沿用已有任务面板的进度、取消、重试、等待连接和人工验证;每页与检查点同事务保存。只有完整分页成功才发布新集合,失败或取消继续使用上一版;首次未完成时不可准备输入。异常字段归属、覆盖率单位或分页协议会失败,不以部分字段代替全集。
|
||||
数据准备支持新建、修改名称和备注、复制、删除及批量删除,详情可查询、添加和移除字段并查看数据集归属。集合固定 `instrument_type + Region + Universe + Delay`,可包含多个同范围数据集;跨范围添加整批拒绝,重复字段去重。空集合可编辑但不能用于研究。
|
||||
|
||||
增量迁移 `0003` 只增加目录、集合、备注与输入表,不改写旧迁移。`/api/v1/catalog` 提供带会话和来源校验的目录/字段查询、完整集合成员、备注、同步创建和输入草稿接口;创建目录同步返回任务 ID,查询、取消及重试仍使用 `/api/v1/sync-jobs`。新任务 `payload` 显式记录范围及数据集,保留旧 Alpha 任务契约。
|
||||
模板工坊、Alpha 变体、特征工程、回测研究、研究流水线和 QuantFlow 共用集合选择弹窗。选择保留集合 ID 和版本,提交研究时核对版本并固定完整字段快照;集合修改、删除或目录重同步均不改变已有研究。AI 可查询集合并固定输入;MCP 用 `search_data_preparations`、`get_data_preparation` 预览,以 `submit_backtests.preparation_refs` 提交按版本选择的集合。
|
||||
|
||||
真实 WorldQuant 数据集 schema、字段所属数据集信息、0–1 覆盖率单位、范围权限和分页协议尚需只读联调。当前证据来自 HTTP 边界合成数据和隔离 PostgreSQL,不代表已验证真实平台兼容性。
|
||||
主要接口:`/api/v1/catalog/fields`(本地)、`/api/v1/catalog/worldquant/fields`(在线)、`/api/v1/data-preparations`(集合)和 `/api/v1/research/input-snapshots/{id}`(研究快照)。旧 `/catalog/inputs`、`/research/inputs` 和已保存输入入口已移除。迁移 `0015` 新建准备集合及独立快照表、移除旧输入表,不迁移旧研究数据;已有目录同步检查点保留。回退需要恢复升级前备份。
|
||||
|
||||
全量字段同步通过 `python -m app.cli catalog-sync` 入队,由现有单进程执行器处理。1Panel 夜间命令、日志、退出码和重试方法见 [生产部署说明](docs/deployment-gitea.md#5-1panel-夜间全量目录同步)。
|
||||
|
||||
实现验收使用模拟上游和隔离数据库。真实 WorldQuant 字段归属、过滤参数、分页协议、范围权限及 1Panel 调度效果需要单独联调。
|
||||
|
||||
## AI 研究助手
|
||||
|
||||
|
||||
@@ -47,6 +47,8 @@ class PageContext(Contract):
|
||||
"alphas",
|
||||
"account",
|
||||
"datasets",
|
||||
"fields",
|
||||
"preparations",
|
||||
"backtests",
|
||||
"operators",
|
||||
"templates",
|
||||
@@ -62,7 +64,7 @@ class PageContext(Contract):
|
||||
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
|
||||
field_id: str | None = Field(default=None, min_length=1, max_length=200)
|
||||
collection_version: str | None = Field(default=None, min_length=1, max_length=36)
|
||||
template_input_id: str | None = Field(default=None, min_length=1, max_length=36)
|
||||
input_snapshot_id: str | None = Field(default=None, min_length=1, max_length=36)
|
||||
unsaved_field_selection: bool = False
|
||||
backtest_run_id: str | None = Field(default=None, max_length=36)
|
||||
backtest_preview_id: str | None = Field(default=None, max_length=36)
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import Literal
|
||||
|
||||
from pydantic import Field, field_validator, model_validator
|
||||
|
||||
from ..preparations.contracts import PreparationReference
|
||||
from ..schemas import Contract
|
||||
|
||||
|
||||
@@ -48,13 +49,18 @@ class Source(Contract):
|
||||
kind: str = Field(default="manual", min_length=1, max_length=100)
|
||||
reference: str | None = Field(default=None, max_length=200)
|
||||
batch_id: str | None = Field(default=None, max_length=200)
|
||||
template_input_id: str | None = Field(default=None, max_length=200)
|
||||
input_snapshot_ids: list[str] = Field(default_factory=list, max_length=100)
|
||||
input_snapshot_id: str | None = Field(default=None, max_length=200)
|
||||
research_id: str | None = Field(default=None, max_length=200)
|
||||
parent_run_id: str | None = Field(default=None, max_length=36)
|
||||
hypothesis: str | None = Field(default=None, max_length=2000)
|
||||
|
||||
|
||||
|
||||
|
||||
class DraftInput(Contract):
|
||||
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
||||
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
source: Source = Field(default_factory=Source)
|
||||
candidates: list[Candidate] = Field(min_length=1, max_length=10000)
|
||||
|
||||
@@ -106,8 +106,33 @@ class Backtests:
|
||||
"limits": "并发和批大小为本地配置,并非平台可用额度;账户外提交不在预算内",
|
||||
}
|
||||
|
||||
async def bind_preparations(self, body):
|
||||
from ..preparations.service import Preparations
|
||||
from ..research.expressions import analyze
|
||||
if not body.preparation_refs and not body.input_ids:
|
||||
return
|
||||
await Preparations(self.db).bind(body)
|
||||
snapshots = [await Preparations(self.db).snapshot(i) for i in body.input_ids]
|
||||
for candidate in body.candidates:
|
||||
scope = dict(instrument_type=candidate.settings.instrumentType, region=candidate.settings.region,
|
||||
universe=candidate.settings.universe, delay=candidate.settings.delay)
|
||||
if any(s["scope"] != scope for s in snapshots):
|
||||
raise HTTPException(422, "数据准备集合与回测范围不一致")
|
||||
fields = {}
|
||||
for snapshot in snapshots:
|
||||
for field, kind in snapshot["field_types"].items():
|
||||
if field in fields and fields[field] != kind:
|
||||
raise HTTPException(422, "输入字段类型冲突")
|
||||
fields[field] = kind
|
||||
validation = analyze(candidate.expression, fields)
|
||||
if validation["syntax"] or validation["types"]:
|
||||
raise HTTPException(422, ";".join(validation["syntax"] + validation["types"]))
|
||||
body.source.input_snapshot_ids = body.input_ids
|
||||
body.source.input_snapshot_id = body.input_ids[0] if len(body.input_ids) == 1 else None
|
||||
|
||||
async def save_draft(self, body, draft_id=None):
|
||||
data = body.model_dump(mode="json", exclude={"version"})
|
||||
await self.bind_preparations(body)
|
||||
data = body.model_dump(mode="json", exclude={"version", "preparation_refs", "input_ids"})
|
||||
if draft_id:
|
||||
changed = await self.db.execute(
|
||||
update(BacktestDraft)
|
||||
@@ -168,6 +193,7 @@ class Backtests:
|
||||
producer. ai_context separately identifies whoever starts the execution.
|
||||
"""
|
||||
if body.inline:
|
||||
await self.bind_preparations(body.inline)
|
||||
data = body.inline.model_dump(mode="json")
|
||||
if self.ai_context and not preserve_source:
|
||||
data["source"] = {
|
||||
|
||||
@@ -263,9 +263,17 @@ class Business:
|
||||
return {"ok": True, "job_id": job_id}
|
||||
|
||||
async def retry_job(self, job_id):
|
||||
job = await self.db.scalar(select(Job).where(Job.id == job_id).with_for_update())
|
||||
# Match create_job's lock order so retry and a fresh scheduled run share one scope owner.
|
||||
await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
|
||||
job = await self.db.scalar(select(Job).where(Job.id == job_id).with_for_update()
|
||||
.execution_options(populate_existing=True))
|
||||
if not job:
|
||||
raise HTTPException(404, "任务不存在")
|
||||
if job.kind == "catalog_full_sync":
|
||||
active = await self.db.scalars(select(Job).where(Job.kind == job.kind, Job.status.in_(ACTIVE)))
|
||||
for existing in active:
|
||||
if existing.payload == job.payload and (existing.id != job.id or job.status in ("queued", "running")):
|
||||
return JobOutput.model_validate(existing).model_dump(mode="json")
|
||||
if job.status not in (
|
||||
"failed",
|
||||
"cancelled",
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from pydantic import AfterValidator, BaseModel, Field, model_validator
|
||||
from pydantic import AfterValidator, BaseModel, Field
|
||||
|
||||
from ..schemas import Contract
|
||||
|
||||
@@ -54,20 +54,6 @@ class NoteInput(Contract):
|
||||
version: int = Field(ge=1)
|
||||
|
||||
|
||||
class InputPreparation(Contract):
|
||||
scope: Scope
|
||||
dataset_id: str = Field(min_length=1, max_length=200)
|
||||
collection_version: str
|
||||
selection: Literal["all", "explicit"] = "all"
|
||||
excluded_ids: list[str] = Field(default_factory=list, max_length=100000)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def valid_selection(self):
|
||||
if self.selection == "all" and self.excluded_ids:
|
||||
raise ValueError("全部字段不能同时提供排除项")
|
||||
return self
|
||||
|
||||
|
||||
class NoteOutput(BaseModel):
|
||||
note: str
|
||||
version: int
|
||||
@@ -107,18 +93,6 @@ class CatalogPage(BaseModel):
|
||||
field_types: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class InputOutput(BaseModel):
|
||||
id: str
|
||||
status: Literal["draft"] = "draft"
|
||||
scope: Scope
|
||||
dataset_id: str
|
||||
collection_version: str
|
||||
selection: str
|
||||
field_ids: list[str]
|
||||
field_types: dict[str, str | None]
|
||||
created_at: UTCTimestamp
|
||||
|
||||
|
||||
class CollectionOutput(BaseModel):
|
||||
collection_version: str | None
|
||||
field_ids: list[str]
|
||||
|
||||
@@ -12,8 +12,6 @@ from .contracts import (
|
||||
CatalogPage,
|
||||
CollectionOutput,
|
||||
EntryOutput,
|
||||
InputOutput,
|
||||
InputPreparation,
|
||||
NoteInput,
|
||||
NoteOutput,
|
||||
Scope,
|
||||
@@ -76,24 +74,6 @@ async def sync(request: Request, body: CatalogJobInput):
|
||||
return result
|
||||
|
||||
|
||||
@router.post("/inputs", status_code=201, response_model=InputOutput)
|
||||
async def prepare(request: Request, body: InputPreparation):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
return await Catalog(db).prepare(body)
|
||||
|
||||
|
||||
@router.get("/inputs", response_model=list[InputOutput])
|
||||
async def inputs(request: Request, scope: Annotated[Scope, Query()]):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Catalog(db).inputs(scope)
|
||||
|
||||
|
||||
@router.get("/inputs/{input_id}", response_model=InputOutput)
|
||||
async def get_input(request: Request, input_id: str):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Catalog(db).input(input_id)
|
||||
|
||||
|
||||
@router.get("/datasets/{dataset_id}/collection", response_model=CollectionOutput)
|
||||
async def collection(request: Request, dataset_id: str, scope: Annotated[Scope, Query()]):
|
||||
async with request.app.state.sessions() as db:
|
||||
|
||||
@@ -4,7 +4,6 @@ The dataset row serializes collection publication and draft creation on PostgreS
|
||||
No page filters participate in template input selection.
|
||||
"""
|
||||
|
||||
from datetime import timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import HTTPException
|
||||
@@ -18,7 +17,6 @@ from ..models import (
|
||||
CatalogNote,
|
||||
CatalogScope,
|
||||
Job,
|
||||
TemplateInput,
|
||||
now,
|
||||
)
|
||||
from ..schemas import JobOutput
|
||||
@@ -164,13 +162,13 @@ class Catalog:
|
||||
raise HTTPException(409, "研究备注已被修改;当前草稿已保留,请载入最新记录后重新保存")
|
||||
return dict(note=body.note, version=body.version + 1, updated_at=now())
|
||||
|
||||
async def create_job(self, body):
|
||||
async def create_job(self, body, *, full=False):
|
||||
account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
|
||||
if not account.password_encrypted or account.connection_status in ("disconnected", "error"):
|
||||
raise HTTPException(409, "请先连接 WorldQuant")
|
||||
if body.dataset_id:
|
||||
await self.dataset(body.scope, body.dataset_id)
|
||||
kind = "field_sync" if body.dataset_id else "catalog_sync"
|
||||
kind = "catalog_full_sync" if full else "field_sync" if body.dataset_id else "catalog_sync"
|
||||
payload = body.model_dump(mode="json")
|
||||
jobs = (
|
||||
await self.db.scalars(
|
||||
@@ -190,7 +188,7 @@ class Catalog:
|
||||
job = Job(id=str(uuid4()), kind=kind, payload=payload)
|
||||
self.db.add(job)
|
||||
await self.db.flush()
|
||||
self.db.add(CatalogBatch(id=job.id, scope_key=body.scope.key(), dataset_id=body.dataset_id))
|
||||
self.db.add(CatalogBatch(id=job.id, job_id=job.id, scope_key=body.scope.key(), dataset_id=body.dataset_id))
|
||||
await self.db.flush()
|
||||
return JobOutput.model_validate(job)
|
||||
|
||||
@@ -210,64 +208,6 @@ class Catalog:
|
||||
)
|
||||
return dict(collection_version=dataset.field_version, field_ids=ids)
|
||||
|
||||
async def prepare(self, body):
|
||||
dataset = await self.dataset(body.scope, body.dataset_id, lock=True)
|
||||
if not dataset.field_version or dataset.field_version != body.collection_version:
|
||||
raise HTTPException(409, "字段集合未完成或版本已变化,请重新读取后准备输入")
|
||||
batch = await self.db.get(CatalogBatch, dataset.field_version)
|
||||
if not batch.complete or batch.scope_key != body.scope.key() or batch.dataset_id != body.dataset_id:
|
||||
raise HTTPException(409, "字段集合不完整")
|
||||
entries = (
|
||||
await self.db.scalars(
|
||||
select(CatalogEntry).where(CatalogEntry.batch_id == batch.id).order_by(CatalogEntry.id)
|
||||
)
|
||||
).all()
|
||||
fields = {e.id: e.field_type for e in entries}
|
||||
excluded = set(body.excluded_ids)
|
||||
if excluded - fields.keys():
|
||||
raise HTTPException(422, "排除项含未知、跨范围或其他数据集字段")
|
||||
chosen = {key: value for key, value in fields.items() if key not in excluded}
|
||||
if not chosen:
|
||||
raise HTTPException(422, "模板输入至少需要一个字段")
|
||||
row = TemplateInput(
|
||||
id=str(uuid4()),
|
||||
scope_key=body.scope.key(),
|
||||
dataset_id=body.dataset_id,
|
||||
collection_version=batch.id,
|
||||
selection=body.selection,
|
||||
field_ids=list(chosen),
|
||||
field_types=chosen,
|
||||
)
|
||||
self.db.add(row)
|
||||
await self.db.flush()
|
||||
return await self.input(row.id)
|
||||
|
||||
async def input(self, input_id):
|
||||
row = await self.db.get(TemplateInput, input_id)
|
||||
if not row:
|
||||
raise HTTPException(404, "输入草稿不存在")
|
||||
scope = await self.db.get(CatalogScope, row.scope_key)
|
||||
return dict(
|
||||
id=row.id,
|
||||
status="draft",
|
||||
scope=scope.scope,
|
||||
dataset_id=row.dataset_id,
|
||||
collection_version=row.collection_version,
|
||||
selection=row.selection,
|
||||
field_ids=row.field_ids,
|
||||
field_types=row.field_types,
|
||||
created_at=row.created_at.replace(tzinfo=timezone.utc)
|
||||
if row.created_at.tzinfo is None
|
||||
else row.created_at,
|
||||
)
|
||||
|
||||
async def inputs(self, scope):
|
||||
ids = (
|
||||
await self.db.scalars(
|
||||
select(TemplateInput.id)
|
||||
.where(TemplateInput.scope_key == scope.key())
|
||||
.order_by(TemplateInput.created_at.desc())
|
||||
.limit(100)
|
||||
)
|
||||
).all()
|
||||
return [await self.input(i) for i in ids]
|
||||
from ..preparations.service import Preparations
|
||||
return await Preparations(self.db).snapshot(input_id)
|
||||
|
||||
+79
-11
@@ -4,6 +4,7 @@ import asyncio
|
||||
import math
|
||||
import re
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
@@ -64,14 +65,15 @@ def normalize(raw, dataset_id):
|
||||
)
|
||||
|
||||
|
||||
async def sync_catalog(runner, job_id, payload):
|
||||
async def sync_catalog(runner, job_id, payload, *, batch_id=None, full=False):
|
||||
scope = Scope.model_validate(payload["scope"])
|
||||
dataset_id = payload.get("dataset_id")
|
||||
batch_id = batch_id or job_id
|
||||
async with runner.sessions() as db:
|
||||
checkpoint = (await db.get(Job, job_id)).checkpoint
|
||||
if checkpoint.get("done"):
|
||||
batch = await db.get(CatalogBatch, batch_id)
|
||||
if batch.complete:
|
||||
return
|
||||
offset = checkpoint.get("offset", 0)
|
||||
offset = batch.offset
|
||||
while True:
|
||||
await runner.checkpoint(job_id, {"next_retry_at": None})
|
||||
raw = await runner.client.catalog_page(scope.model_dump(), dataset_id, offset)
|
||||
@@ -97,12 +99,12 @@ async def sync_catalog(runner, job_id, payload):
|
||||
job = await db.get(Job, job_id)
|
||||
if job.cancel_requested:
|
||||
raise asyncio.CancelledError()
|
||||
batch = await db.get(CatalogBatch, job_id)
|
||||
batch = await db.get(CatalogBatch, batch_id)
|
||||
added = 0
|
||||
for entry in entries:
|
||||
if await db.get(CatalogEntry, (job_id, entry["id"])):
|
||||
if await db.get(CatalogEntry, (batch_id, entry["id"])):
|
||||
continue
|
||||
db.add(CatalogEntry(batch_id=job_id, **entry))
|
||||
db.add(CatalogEntry(batch_id=batch_id, **entry))
|
||||
await db.flush()
|
||||
added += 1
|
||||
owner = dataset_id or entry["id"]
|
||||
@@ -112,12 +114,15 @@ async def sync_catalog(runner, job_id, payload):
|
||||
if rows and not added:
|
||||
raise WqError("平台分页重复且未前进,已保留进度", "invalid_response")
|
||||
batch.count += added
|
||||
if not full:
|
||||
job.processed = batch.count
|
||||
offset += len(rows)
|
||||
job.checkpoint = dict(offset=offset, done=not more)
|
||||
batch.offset = offset
|
||||
job.checkpoint = {**job.checkpoint, "offset": offset, "done": not more, "current_field_count": batch.count}
|
||||
job.updated_at = now()
|
||||
if not more:
|
||||
batch.complete, batch.completed_at = True, now()
|
||||
if not full:
|
||||
job.total = batch.count
|
||||
if dataset_id:
|
||||
dataset = await db.scalar(
|
||||
@@ -125,12 +130,12 @@ async def sync_catalog(runner, job_id, payload):
|
||||
.where(CatalogDataset.scope_key == scope.key(), CatalogDataset.id == dataset_id)
|
||||
.with_for_update()
|
||||
)
|
||||
dataset.field_version = job_id
|
||||
dataset.field_version = batch_id
|
||||
else:
|
||||
scope_row = await db.get(CatalogScope, scope.key())
|
||||
scope_row.catalog_version, scope_row.synced_at = job_id, now()
|
||||
scope_row.catalog_version, scope_row.synced_at = batch_id, now()
|
||||
ids = (
|
||||
await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id))
|
||||
await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == batch_id))
|
||||
).all()
|
||||
for item_id in ids:
|
||||
if not await db.get(CatalogDataset, (scope.key(), item_id)):
|
||||
@@ -138,3 +143,66 @@ async def sync_catalog(runner, job_id, payload):
|
||||
await db.commit()
|
||||
if not more:
|
||||
return
|
||||
|
||||
|
||||
async def sync_full_catalog(runner, job_id, payload):
|
||||
"""Resume each dataset batch independently; only publish complete enumerations."""
|
||||
scope = Scope.model_validate(payload["scope"])
|
||||
options = await runner.client.get_platform_setting_options()
|
||||
if not any(r["instrument_type"] == scope.instrument_type and r["region"] == scope.region
|
||||
and r["delay"] == scope.delay and scope.universe in r["universes"]
|
||||
for r in options["instrument_options"]):
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
job.checkpoint = {**job.checkpoint, "error_code": "invalid_scope"}
|
||||
raise WqError("平台不支持该研究范围", "invalid_scope")
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
job.checkpoint = {**{k: v for k, v in job.checkpoint.items() if k != "error_code"}, "phase": "catalog"}
|
||||
await sync_catalog(runner, job_id, {"scope": payload["scope"]}, full=True)
|
||||
async with runner.sessions() as db:
|
||||
ids = list(await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id)
|
||||
.order_by(CatalogEntry.id)))
|
||||
job = await db.get(Job, job_id)
|
||||
failures = dict(job.checkpoint.get("failures", {}))
|
||||
completed = 0
|
||||
await runner.checkpoint(job_id, {"total": len(ids)})
|
||||
for dataset_id in ids:
|
||||
async with runner.sessions.begin() as db:
|
||||
batch = await db.scalar(select(CatalogBatch).where(CatalogBatch.job_id == job_id,
|
||||
CatalogBatch.dataset_id == dataset_id))
|
||||
if not batch:
|
||||
batch = CatalogBatch(id=str(uuid4()), job_id=job_id, scope_key=scope.key(), dataset_id=dataset_id)
|
||||
db.add(batch)
|
||||
await db.flush()
|
||||
batch_id, complete = batch.id, batch.complete
|
||||
job = await db.get(Job, job_id)
|
||||
if job.cancel_requested:
|
||||
raise asyncio.CancelledError()
|
||||
job.checkpoint = {**job.checkpoint, "phase": "fields", "dataset_id": dataset_id,
|
||||
"datasets_completed": completed, "datasets_total": len(ids),
|
||||
"offset": batch.offset, "current_field_count": batch.count}
|
||||
if not complete:
|
||||
try:
|
||||
await sync_catalog(runner, job_id, {"scope": payload["scope"], "dataset_id": dataset_id},
|
||||
batch_id=batch_id, full=True)
|
||||
except WqError as exc:
|
||||
if exc.code in ("disconnected", "authentication_failed", "identity_mismatch", "verification_required", "network_error"):
|
||||
raise
|
||||
failures[dataset_id] = str(exc)
|
||||
if complete or (await _batch_complete(runner, batch_id)):
|
||||
failures.pop(dataset_id, None)
|
||||
completed += 1
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
job.processed, job.failed = completed, len(failures)
|
||||
job.checkpoint = {**job.checkpoint, "datasets_completed": completed, "failures": failures}
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
job.error = f"{len(failures)} 个数据集同步失败" if failures else None
|
||||
job.checkpoint = {**job.checkpoint, "phase": "finished"}
|
||||
|
||||
|
||||
async def _batch_complete(runner, batch_id):
|
||||
async with runner.sessions() as db:
|
||||
return (await db.get(CatalogBatch, batch_id)).complete
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import getpass
|
||||
import math
|
||||
|
||||
from sqlalchemy import delete, update
|
||||
|
||||
@@ -59,6 +60,63 @@ async def token_command(args):
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def catalog_sync_command(args):
|
||||
"""Enqueue on the existing runner, then observe without owning the upstream session."""
|
||||
import json
|
||||
import time
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from .business import Business
|
||||
from .catalog.contracts import CatalogJobInput, Scope
|
||||
from .catalog.service import Catalog
|
||||
from .models import Job
|
||||
|
||||
engine, sessions = create_database(Settings().database_url)
|
||||
try:
|
||||
async with sessions.begin() as db:
|
||||
if args.resume_job:
|
||||
if args.region or args.universe or args.delay is not None:
|
||||
raise ValueError("--resume-job 不能同时指定新范围")
|
||||
job = await db.get(Job, args.resume_job)
|
||||
if not job or job.kind != "catalog_full_sync":
|
||||
raise ValueError("只能恢复已有全量目录任务")
|
||||
result = await Business(db).retry_job(job.id)
|
||||
job_id = result["id"]
|
||||
else:
|
||||
if not args.region or not args.universe or args.delay is None:
|
||||
raise ValueError("需要 --region、--universe 和 --delay")
|
||||
scope = Scope(instrument_type=args.instrument_type, region=args.region,
|
||||
universe=args.universe, delay=args.delay)
|
||||
job = await Catalog(db).create_job(CatalogJobInput(scope=scope), full=True)
|
||||
job_id = job.id
|
||||
deadline, previous = time.monotonic() + args.wait_timeout, None
|
||||
while True:
|
||||
async with sessions() as db:
|
||||
job = await db.get(Job, job_id)
|
||||
data = dict(job_id=job.id, status=job.status, processed=job.processed,
|
||||
total=job.total, failed=job.failed, checkpoint=job.checkpoint, error=job.error)
|
||||
current = json.dumps(data, ensure_ascii=False, sort_keys=True)
|
||||
if current != previous:
|
||||
print(current, flush=True)
|
||||
previous = current
|
||||
if job.status == "completed":
|
||||
return 0
|
||||
if job.status in ("failed", "completed_with_errors", "cancelled"):
|
||||
return 2 if job.checkpoint.get("error_code") == "invalid_scope" else 1
|
||||
if job.status in ("waiting_auth", "waiting_connection"):
|
||||
return 3
|
||||
if time.monotonic() >= deadline:
|
||||
print(f"等待超时;后台任务 {job_id} 继续执行", flush=True)
|
||||
return 4
|
||||
await asyncio.sleep(min(5, max(0, deadline - time.monotonic())))
|
||||
except HTTPException as exc:
|
||||
print(str(exc.detail), flush=True)
|
||||
return 3 if exc.status_code == 409 else 1
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
commands = parser.add_subparsers(dest="command", required=True)
|
||||
@@ -70,8 +128,19 @@ if __name__ == "__main__":
|
||||
commands.add_parser("mcp-token-list")
|
||||
revoke = commands.add_parser("mcp-token-revoke")
|
||||
revoke.add_argument("token_id")
|
||||
sync = commands.add_parser("catalog-sync", help="全量同步一个范围的数据集及全部字段")
|
||||
sync.add_argument("--region")
|
||||
sync.add_argument("--universe")
|
||||
sync.add_argument("--delay", type=int, choices=range(0, 10))
|
||||
sync.add_argument("--instrument-type", default="EQUITY")
|
||||
sync.add_argument("--resume-job")
|
||||
sync.add_argument("--wait-timeout", type=float, default=21600)
|
||||
args = parser.parse_args()
|
||||
if args.command == "catalog-sync" and (not math.isfinite(args.wait_timeout) or args.wait_timeout <= 0):
|
||||
parser.error("--wait-timeout 必须大于 0")
|
||||
try:
|
||||
if args.command == "catalog-sync":
|
||||
raise SystemExit(asyncio.run(catalog_sync_command(args)))
|
||||
asyncio.run(reset_password() if args.command == "reset-password" else token_command(args))
|
||||
except ValueError as exc:
|
||||
parser.error(str(exc))
|
||||
|
||||
+5
-1
@@ -255,6 +255,10 @@ class Runner:
|
||||
await self.ensure_connected(force=kind == "connect")
|
||||
if kind in ("connect", "profile"):
|
||||
await self.refresh_profile()
|
||||
elif kind == "catalog_full_sync":
|
||||
from .catalog.sync import sync_full_catalog
|
||||
|
||||
await sync_full_catalog(self, job_id, payload)
|
||||
elif kind in ("catalog_sync", "field_sync"):
|
||||
from .catalog.sync import sync_catalog
|
||||
|
||||
@@ -298,7 +302,7 @@ class Runner:
|
||||
await self.checkpoint(
|
||||
job_id,
|
||||
{
|
||||
"status": "waiting_connection" if waiting else "failed",
|
||||
"status": "waiting_connection" if waiting or (kind == "catalog_full_sync" and exc.code == "network_error") else "failed",
|
||||
"error": str(exc),
|
||||
"next_retry_at": None,
|
||||
},
|
||||
|
||||
@@ -27,6 +27,7 @@ from .db import create_database
|
||||
from .jobs import AUTH_KINDS, Runner, create_job
|
||||
from .mcp_api.token_routes import router as mcp_token_router
|
||||
from .models import Account, Admin, BacktestConfig, Job, JobItem, LoginSession
|
||||
from .preparations.routes import router as preparations_router
|
||||
from .research.routes import router as research_router
|
||||
from .research.runtime import ResearchRuntime
|
||||
from .schemas import (
|
||||
@@ -484,6 +485,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
app.include_router(backtest_router)
|
||||
app.include_router(api)
|
||||
app.include_router(catalog_router)
|
||||
app.include_router(preparations_router)
|
||||
app.include_router(research_catalog_router)
|
||||
app.include_router(research_router)
|
||||
app.include_router(ai_router(ai_runtime))
|
||||
|
||||
@@ -21,6 +21,8 @@ from ..research_access.service import ResearchAccess, ResearchError
|
||||
|
||||
# Name, schema, business method, required scope, description. No generic arbitrary HTTP tool.
|
||||
TOOLS = {
|
||||
"search_data_preparations": (c.PreparationSearch, "preparations", "research:read", "分页查询数据准备集合,返回固定范围、字段数及版本。研究可使用多个集合,各集合范围独立。"),
|
||||
"get_data_preparation": (c.PreparationRead, "preparation", "research:read", "按集合 ID 与版本分页预览字段、类型、描述和数据集归属。提交回测时携带 preparation_refs,由服务端核对版本并固定独立输入快照;空集合不能用于研究。"),
|
||||
"create_research_template": (c.CreateTemplate, "create_template", "research:write", "将调用方大模型研究后自行总结的参数化模板保存到模板工坊,供用户后续批量回测。先用 get_backtest_results 阅读实际指标和检查,选择 1–20 个已完成采集的 source_item_ids,并说明 hypothesis;不要把 completed 当作检查通过。template 使用 {name} 占位符及逐一对应的 variables,字段变量须声明 MATRIX/VECTOR/GROUP,VECTOR 聚合须明确写入表达式。提供唯一名称和 idempotency_key,可附 reference。返回模板 ID、版本和理论组合数;仅核验结构及来源,不验证所有参数组合,不再次调用模型、不执行回测、不覆盖已有模板。"),
|
||||
"get_submission_check": (c.SelfCorrelationReference, "submission_check_context", "research:read", "读取已导入 Alpha 的表达式、Description、snapshot 和缓存检查结果;不发起检查。先核对或生成三段 Description,再调用 check_submission。"),
|
||||
"check_submission": (c.SubmissionCheck, "check_submission", "research:refresh", "对单个待提交 Alpha 写回已获用户授权的 Description 并调用平台 GET /check,返回 job_id。须先用 get_submission_check 获取 snapshot;保留本地自相关门槛和冲突保护。通过 get_refresh_job 查进度、get_submission_check 读结果。无论检查结果如何,都不会调用 /submit 或正式提交 Alpha。"),
|
||||
@@ -34,7 +36,7 @@ TOOLS = {
|
||||
"check_self_correlation": (c.SelfCorrelationCheck, "check_self_correlation", "research:refresh", "对 1–100 个已导入 Alpha 发起本地自相关检查,返回 job_id。与本地同地区已提交 Alpha 比较,排除自身;优先用缓存,缺失 PnL 自动补取。需先同步已提交 Alpha;不调用平台提交检查。用 get_refresh_job 查进度、get_self_correlation 读结果。"),
|
||||
"get_self_correlation": (c.SelfCorrelationReference, "self_correlation", "research:read", "读取指定 Alpha 最新的本地自相关缓存,包括最大相关系数、样本覆盖与 stale 状态;无缓存不自动检查,结果不等同于平台提交资格。"),
|
||||
"search_backtests": (c.History, "history", "research:read", "分页查历史候选与固定设置;candidates 按完整输入精确匹配,不推断数学等价。"),
|
||||
"submit_backtests": (c.Submit, "submit", "backtests:execute", "执行用户已授权的固定批次,自动留痕并立即返回运行 ID。每项必须完整设置;重复默认拒绝,rerun 明确重跑。不需要研究资产。"),
|
||||
"submit_backtests": (c.Submit, "submit", "backtests:execute", "执行用户已授权的固定批次,自动留痕并立即返回运行 ID。可携带 preparation_refs 选择集合,版本变化须重新读取;每项必须完整设置;重复默认拒绝,rerun 明确重跑。不需要研究资产。"),
|
||||
"get_backtest": (c.RunReference, "run", "research:read", "读取真实运行进度、提交数量和可选增量事件;受理不等于成功。"),
|
||||
"get_backtest_results": (c.Results, "results", "research:read", "分页读取固定快照指标、全部非通过检查及三层状态;缺失指标不补零。"),
|
||||
"get_backtest_artifact": (c.Artifact, "artifact", "research:read", "分页读取候选脱敏快照的顶层键值或独立采集的 PnL;缺缓存不自动刷新。"),
|
||||
|
||||
+31
-9
@@ -340,7 +340,9 @@ class CatalogScope(Base):
|
||||
|
||||
class CatalogBatch(Base):
|
||||
__tablename__ = "catalog_batches"
|
||||
id: Mapped[str] = mapped_column(ForeignKey("sync_jobs.id"), primary_key=True)
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True)
|
||||
job_id: Mapped[str] = mapped_column(ForeignKey("sync_jobs.id"), index=True)
|
||||
offset: Mapped[int] = mapped_column(Integer, default=0)
|
||||
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True)
|
||||
dataset_id: Mapped[str | None] = mapped_column(String(200))
|
||||
complete: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
@@ -386,16 +388,36 @@ class CatalogNote(Base):
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class TemplateInput(Base):
|
||||
__tablename__ = "template_inputs"
|
||||
class DataPreparation(Base):
|
||||
"""Editable collection; scope never changes after creation."""
|
||||
__tablename__ = "data_preparations"
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True)
|
||||
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True)
|
||||
dataset_id: Mapped[str] = mapped_column(String(200))
|
||||
collection_version: Mapped[str] = mapped_column(ForeignKey("catalog_batches.id"))
|
||||
selection: Mapped[str] = mapped_column(String(20))
|
||||
field_ids: Mapped[list] = mapped_column(JSON)
|
||||
field_types: Mapped[dict] = mapped_column(JSON)
|
||||
name: Mapped[str] = mapped_column(String(200))
|
||||
note: Mapped[str] = mapped_column(Text, default="")
|
||||
scope_key: Mapped[str] = mapped_column(String(200), index=True)
|
||||
scope: Mapped[dict] = mapped_column(JSON)
|
||||
version: Mapped[int] = mapped_column(Integer, default=1)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class PreparationField(Base):
|
||||
__tablename__ = "preparation_fields"
|
||||
preparation_id: Mapped[str] = mapped_column(ForeignKey("data_preparations.id", ondelete="CASCADE"), primary_key=True)
|
||||
field_id: Mapped[str] = mapped_column(String(200), primary_key=True)
|
||||
dataset_id: Mapped[str] = mapped_column(String(200), index=True)
|
||||
content: Mapped[dict] = mapped_column(JSON)
|
||||
|
||||
|
||||
class ResearchInputSnapshot(Base):
|
||||
"""Self-contained research input: deletion of its preparation cannot invalidate it."""
|
||||
__tablename__ = "research_input_snapshots"
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True)
|
||||
preparation_id: Mapped[str] = mapped_column(String(36), index=True)
|
||||
preparation_version: Mapped[int] = mapped_column(Integer)
|
||||
content: Mapped[dict] = mapped_column(JSON)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
__table_args__ = (UniqueConstraint("preparation_id", "preparation_version"),)
|
||||
|
||||
|
||||
class CatalogResource(Base):
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Data preparation collections and immutable research inputs."""
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Shared collection and field-query contracts."""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import Field, model_validator
|
||||
|
||||
from ..catalog.contracts import Scope
|
||||
from ..schemas import Contract
|
||||
|
||||
|
||||
class FieldFilters(Scope):
|
||||
q: str = Field(default="", max_length=300)
|
||||
dataset_id: str | None = None
|
||||
category: str | None = None
|
||||
subcategory: str | None = None
|
||||
field_type: str | None = None
|
||||
coverage_min: float | None = Field(default=None, ge=0, le=1)
|
||||
coverage_max: float | None = Field(default=None, ge=0, le=1)
|
||||
user_count_min: int | None = Field(default=None, ge=0)
|
||||
user_count_max: int | None = Field(default=None, ge=0)
|
||||
alpha_count_min: int | None = Field(default=None, ge=0)
|
||||
alpha_count_max: int | None = Field(default=None, ge=0)
|
||||
synced_from: datetime | None = None
|
||||
synced_to: datetime | None = None
|
||||
sort: Literal["id", "name", "dataset_id", "coverage", "user_count", "alpha_count", "synced_at"] = "id"
|
||||
direction: Literal["asc", "desc"] = "asc"
|
||||
limit: int = Field(default=25, ge=1, le=100)
|
||||
offset: int = Field(default=0, ge=0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def ranges(self):
|
||||
for key in ("coverage", "user_count", "alpha_count"):
|
||||
low, high = getattr(self, key + "_min"), getattr(self, key + "_max")
|
||||
if low is not None and high is not None and low > high:
|
||||
raise ValueError("筛选下限不能超过上限")
|
||||
return self
|
||||
|
||||
|
||||
class FieldReference(Contract):
|
||||
scope: Scope
|
||||
dataset_id: str = Field(min_length=1, max_length=200)
|
||||
field_id: str = Field(min_length=1, max_length=200)
|
||||
source: Literal["local", "worldquant"] = "local"
|
||||
collection_version: str | None = None
|
||||
|
||||
|
||||
class PreparationCreate(Contract):
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
note: str = Field(default="", max_length=20000)
|
||||
scope: Scope
|
||||
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
|
||||
|
||||
|
||||
class PreparationVersion(Contract):
|
||||
version: int = Field(ge=1)
|
||||
|
||||
|
||||
class PreparationEdit(PreparationVersion):
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
note: str = Field(default="", max_length=20000)
|
||||
|
||||
|
||||
class MemberChange(PreparationVersion):
|
||||
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
|
||||
remove_ids: list[str] = Field(default_factory=list, max_length=10000)
|
||||
|
||||
|
||||
class PreparationReference(PreparationVersion):
|
||||
id: str = Field(min_length=1, max_length=36)
|
||||
|
||||
|
||||
class PreparationReferences(Contract):
|
||||
items: list[PreparationReference] = Field(min_length=1, max_length=100)
|
||||
|
||||
|
||||
class DatasetCopy(Contract):
|
||||
scope: Scope
|
||||
dataset_id: str = Field(min_length=1, max_length=200)
|
||||
collection_version: str
|
||||
@@ -0,0 +1,195 @@
|
||||
"""Authenticated preparation and field-directory endpoints."""
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import delete, select
|
||||
|
||||
from ..catalog.contracts import Scope
|
||||
from ..models import CatalogScope, PreparationField, now
|
||||
from ..security import require_auth
|
||||
from .contracts import (
|
||||
DatasetCopy,
|
||||
FieldFilters,
|
||||
MemberChange,
|
||||
PreparationCreate,
|
||||
PreparationEdit,
|
||||
PreparationReferences,
|
||||
PreparationVersion,
|
||||
)
|
||||
from .service import Preparations
|
||||
|
||||
router = APIRouter(prefix="/api/v1", dependencies=[Depends(require_auth)], tags=["data-preparations"])
|
||||
|
||||
|
||||
@router.get("/catalog/local-scopes")
|
||||
async def local_scopes(request: Request):
|
||||
async with request.app.state.sessions() as db:
|
||||
rows = await db.scalars(select(CatalogScope))
|
||||
options = {}
|
||||
for row in rows:
|
||||
key = (row.scope["instrument_type"], row.scope["region"], row.scope["delay"])
|
||||
option = options.setdefault(
|
||||
key, {k: row.scope[k] for k in ("instrument_type", "region", "delay")}
|
||||
)
|
||||
option.setdefault("universes", []).append(row.scope["universe"])
|
||||
return {"instrument_options": list(options.values())}
|
||||
|
||||
|
||||
@router.get("/catalog/fields")
|
||||
async def local_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Preparations(db).fields(filters)
|
||||
|
||||
|
||||
@router.get("/catalog/worldquant/fields")
|
||||
async def online_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Preparations(db, request.app.state.runner.client).online_fields(filters)
|
||||
|
||||
|
||||
@router.get("/data-preparations")
|
||||
async def preparations(
|
||||
request: Request,
|
||||
q: str = "",
|
||||
scope_key: str | None = None,
|
||||
limit: int = Query(25, ge=1, le=100),
|
||||
offset: int = Query(0, ge=0),
|
||||
):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Preparations(db).list(q, scope_key, limit, offset)
|
||||
|
||||
|
||||
@router.post("/data-preparations", status_code=201)
|
||||
async def create(request: Request, body: PreparationCreate):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
service = Preparations(db, request.app.state.runner.client)
|
||||
fields = await service.resolve_fields(body.scope, body.fields)
|
||||
return await service.create(body.name, body.note, body.scope, fields)
|
||||
|
||||
|
||||
@router.post("/data-preparations/from-dataset", status_code=201)
|
||||
async def from_dataset(request: Request, body: DatasetCopy):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
return await Preparations(db).copy_dataset(body)
|
||||
|
||||
|
||||
@router.post("/data-preparations/batch-delete")
|
||||
async def batch_delete(request: Request, body: PreparationReferences):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
return await Preparations(db).remove(body.items)
|
||||
|
||||
|
||||
@router.post("/data-preparations/freeze", status_code=201)
|
||||
async def freeze(request: Request, body: PreparationReferences):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
return {"items": await Preparations(db).freeze(body.items)}
|
||||
|
||||
|
||||
@router.get("/research/input-snapshots/{snapshot_id}")
|
||||
async def snapshot(request: Request, snapshot_id: str):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Preparations(db).snapshot(snapshot_id)
|
||||
|
||||
|
||||
@router.get("/data-preparations/{preparation_id}")
|
||||
async def detail(request: Request, preparation_id: str):
|
||||
async with request.app.state.sessions() as db:
|
||||
service = Preparations(db)
|
||||
return await service.output(await service.get(preparation_id))
|
||||
|
||||
|
||||
@router.patch("/data-preparations/{preparation_id}")
|
||||
async def edit(request: Request, preparation_id: str, body: PreparationEdit):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
service = Preparations(db)
|
||||
row = await service.get(preparation_id, body.version, lock=True)
|
||||
row.name, row.note, row.updated_at, row.version = body.name, body.note, now(), row.version + 1
|
||||
return await service.output(row)
|
||||
|
||||
|
||||
@router.delete("/data-preparations/{preparation_id}")
|
||||
async def remove(request: Request, preparation_id: str, version: int = Query(ge=1)):
|
||||
from .contracts import PreparationReference
|
||||
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
return await Preparations(db).remove([PreparationReference(id=preparation_id, version=version)])
|
||||
|
||||
|
||||
@router.post("/data-preparations/{preparation_id}/copy", status_code=201)
|
||||
async def copy(request: Request, preparation_id: str, body: PreparationVersion):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
service = Preparations(db)
|
||||
row = await service.get(preparation_id, body.version, lock=True)
|
||||
fields = [
|
||||
f.content
|
||||
for f in await db.scalars(
|
||||
select(PreparationField).where(PreparationField.preparation_id == row.id)
|
||||
)
|
||||
]
|
||||
return await service.create(
|
||||
(row.name + " 副本")[:200], row.note, Scope.model_validate(row.scope), fields
|
||||
)
|
||||
|
||||
|
||||
@router.get("/data-preparations/{preparation_id}/fields")
|
||||
async def members(
|
||||
request: Request,
|
||||
preparation_id: str,
|
||||
q: str = "",
|
||||
dataset_id: str | None = None,
|
||||
limit: int = Query(25, ge=1, le=100),
|
||||
offset: int = Query(0, ge=0),
|
||||
):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Preparations(db).members(preparation_id, q, dataset_id, limit, offset)
|
||||
|
||||
|
||||
@router.patch("/data-preparations/{preparation_id}/fields")
|
||||
async def change_members(request: Request, preparation_id: str, body: MemberChange):
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
service = Preparations(db, request.app.state.runner.client)
|
||||
row = await service.get(preparation_id, body.version, lock=True)
|
||||
fields = await service.resolve_fields(Scope.model_validate(row.scope), body.fields)
|
||||
await service.add(row, fields)
|
||||
if body.remove_ids:
|
||||
present = set(
|
||||
await db.scalars(
|
||||
select(PreparationField.field_id).where(
|
||||
PreparationField.preparation_id == row.id,
|
||||
PreparationField.field_id.in_(body.remove_ids),
|
||||
)
|
||||
)
|
||||
)
|
||||
if set(body.remove_ids) - present:
|
||||
raise HTTPException(422, "移除项含不属于该集合的字段")
|
||||
await db.execute(
|
||||
delete(PreparationField).where(
|
||||
PreparationField.preparation_id == row.id, PreparationField.field_id.in_(body.remove_ids)
|
||||
)
|
||||
)
|
||||
row.version, row.updated_at = row.version + 1, now()
|
||||
return await service.output(row)
|
||||
|
||||
|
||||
@router.get("/data-preparations/{preparation_id}/selection")
|
||||
async def selection(request: Request, preparation_id: str, version: int = Query(ge=1)):
|
||||
async with request.app.state.sessions() as db:
|
||||
service = Preparations(db)
|
||||
row = await service.get(preparation_id, version)
|
||||
fields = [
|
||||
f.content
|
||||
for f in await db.scalars(
|
||||
select(PreparationField)
|
||||
.where(PreparationField.preparation_id == row.id)
|
||||
.order_by(PreparationField.field_id)
|
||||
)
|
||||
]
|
||||
return {
|
||||
**await service.output(row),
|
||||
"fields": fields,
|
||||
"field_ids": [f["id"] for f in fields],
|
||||
"field_types": {f["id"]: f["field_type"] for f in fields},
|
||||
"dataset_ids": sorted({f["dataset_id"] for f in fields}),
|
||||
"preparation_ref": {"id": row.id, "version": row.version},
|
||||
}
|
||||
@@ -0,0 +1,450 @@
|
||||
"""Collection operations own validation; callers own transactions and authorization."""
|
||||
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import delete, func, or_, select
|
||||
from sqlalchemy.orm import aliased
|
||||
|
||||
from ..catalog.contracts import EntryOutput, Scope
|
||||
from ..catalog.research_metadata import upstream
|
||||
from ..catalog.service import Catalog
|
||||
from ..catalog.sync import identifier, label, normalize
|
||||
from ..models import (
|
||||
CatalogBatch,
|
||||
CatalogDataset,
|
||||
CatalogEntry,
|
||||
CatalogScope,
|
||||
DataPreparation,
|
||||
PreparationField,
|
||||
ResearchInputSnapshot,
|
||||
now,
|
||||
)
|
||||
from ..research.serialization import encode_snapshot
|
||||
from .contracts import FieldFilters
|
||||
|
||||
|
||||
def page(items, total, limit, offset, **extra):
|
||||
return dict(
|
||||
items=items, total=total, limit=limit, offset=offset, has_more=offset + len(items) < total, **extra
|
||||
)
|
||||
|
||||
|
||||
def contains(value):
|
||||
return "%" + value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
|
||||
|
||||
|
||||
class Preparations:
|
||||
def __init__(self, db, client=None):
|
||||
self.db, self.client = db, client
|
||||
|
||||
async def fields(self, filters):
|
||||
scope = await self.db.get(CatalogScope, filters.key())
|
||||
owner = aliased(CatalogEntry)
|
||||
query = (
|
||||
select(
|
||||
CatalogEntry,
|
||||
CatalogBatch.dataset_id,
|
||||
owner.name.label("dataset_name"),
|
||||
owner.category,
|
||||
owner.subcategory,
|
||||
)
|
||||
.join(CatalogBatch, CatalogEntry.batch_id == CatalogBatch.id)
|
||||
.join(
|
||||
CatalogDataset,
|
||||
(CatalogDataset.field_version == CatalogBatch.id)
|
||||
& (CatalogDataset.scope_key == filters.key()),
|
||||
)
|
||||
.outerjoin(
|
||||
owner,
|
||||
(owner.id == CatalogBatch.dataset_id)
|
||||
& (owner.batch_id == (scope.catalog_version if scope else None)),
|
||||
)
|
||||
.where(CatalogBatch.complete.is_(True))
|
||||
)
|
||||
if filters.q:
|
||||
query = query.where(
|
||||
or_(
|
||||
*[
|
||||
column.ilike(contains(filters.q), escape="\\")
|
||||
for column in (
|
||||
CatalogEntry.id,
|
||||
CatalogEntry.name,
|
||||
CatalogEntry.description,
|
||||
CatalogBatch.dataset_id,
|
||||
owner.name,
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
for key in ("dataset_id", "field_type", "category", "subcategory"):
|
||||
value = getattr(filters, key)
|
||||
column = (
|
||||
CatalogBatch.dataset_id
|
||||
if key == "dataset_id"
|
||||
else func.coalesce(getattr(CatalogEntry, key), getattr(owner, key))
|
||||
if key in ("category", "subcategory")
|
||||
else getattr(CatalogEntry, key)
|
||||
)
|
||||
if value:
|
||||
query = query.where(column == value)
|
||||
for key in ("coverage", "user_count", "alpha_count"):
|
||||
low, high = getattr(filters, key + "_min"), getattr(filters, key + "_max")
|
||||
if low is not None:
|
||||
query = query.where(getattr(CatalogEntry, key) >= low)
|
||||
if high is not None:
|
||||
query = query.where(getattr(CatalogEntry, key) <= high)
|
||||
if filters.synced_from:
|
||||
query = query.where(CatalogEntry.synced_at >= filters.synced_from)
|
||||
if filters.synced_to:
|
||||
query = query.where(CatalogEntry.synced_at <= filters.synced_to)
|
||||
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
|
||||
column = (
|
||||
CatalogBatch.dataset_id if filters.sort == "dataset_id" else getattr(CatalogEntry, filters.sort)
|
||||
)
|
||||
query = query.order_by(
|
||||
(column.desc() if filters.direction == "desc" else column.asc()).nulls_last(),
|
||||
CatalogBatch.dataset_id,
|
||||
CatalogEntry.id,
|
||||
)
|
||||
rows = (await self.db.execute(query.limit(filters.limit).offset(filters.offset))).all()
|
||||
items = [
|
||||
{
|
||||
**self.local_field(row, dataset_id, name, filters),
|
||||
"category": row.category or category,
|
||||
"subcategory": row.subcategory or subcategory,
|
||||
}
|
||||
for row, dataset_id, name, category, subcategory in rows
|
||||
]
|
||||
return page(items, total, filters.limit, filters.offset)
|
||||
|
||||
@staticmethod
|
||||
def local_field(row, dataset_id, name, scope):
|
||||
return encode_snapshot(
|
||||
dict(
|
||||
**EntryOutput.model_validate(row, from_attributes=True).model_dump(
|
||||
exclude={"scope", "dataset_id", "collection_version"}
|
||||
),
|
||||
field_id=row.id,
|
||||
dataset_id=dataset_id,
|
||||
dataset_name=name or dataset_id,
|
||||
collection_version=row.batch_id,
|
||||
scope=scope.model_dump(include=set(Scope.model_fields)),
|
||||
source="local",
|
||||
fetched_at=row.synced_at,
|
||||
)
|
||||
)
|
||||
|
||||
async def online_fields(self, filters):
|
||||
if self.client is None:
|
||||
raise HTTPException(409, "请先连接 WorldQuant")
|
||||
params = dict(
|
||||
instrumentType=filters.instrument_type,
|
||||
region=filters.region,
|
||||
universe=filters.universe,
|
||||
delay=filters.delay,
|
||||
limit=filters.limit,
|
||||
offset=filters.offset,
|
||||
)
|
||||
for key, remote in (("q", "search"), ("dataset_id", "dataset.id"), ("field_type", "type")):
|
||||
value = getattr(filters, key)
|
||||
if value:
|
||||
params[remote] = value
|
||||
for key, remote in (
|
||||
("coverage", "coverage"),
|
||||
("user_count", "userCount"),
|
||||
("alpha_count", "alphaCount"),
|
||||
):
|
||||
for suffix, op in (("min", ">"), ("max", "<")):
|
||||
value = getattr(filters, key + "_" + suffix)
|
||||
if value is not None:
|
||||
params[remote + op] = value
|
||||
raw = await upstream(self.client.get("/data-fields", params))
|
||||
rows = raw.get("results")
|
||||
if not isinstance(rows, list):
|
||||
raise HTTPException(502, "平台字段列表格式无法识别")
|
||||
items = []
|
||||
for row in rows:
|
||||
if not isinstance(row, dict):
|
||||
raise HTTPException(502, "平台字段记录格式无法识别")
|
||||
dataset = row.get("dataset")
|
||||
owner = dataset.get("id") if isinstance(dataset, dict) else dataset
|
||||
try:
|
||||
owner = identifier(owner)
|
||||
data = normalize(row, owner)
|
||||
except Exception as exc:
|
||||
from ..worldquant import WqError
|
||||
|
||||
if isinstance(exc, WqError):
|
||||
raise HTTPException(502, str(exc)) from None
|
||||
raise
|
||||
if filters.dataset_id and owner != filters.dataset_id:
|
||||
raise HTTPException(502, "平台返回其他数据集字段")
|
||||
for remote_key, key in (
|
||||
("instrumentType", "instrument_type"),
|
||||
("instrument_type", "instrument_type"),
|
||||
("region", "region"),
|
||||
("universe", "universe"),
|
||||
("delay", "delay"),
|
||||
):
|
||||
if remote_key in row and row[remote_key] != getattr(filters, key):
|
||||
raise HTTPException(502, "平台字段范围与查询不一致")
|
||||
items.append(
|
||||
encode_snapshot(
|
||||
dict(
|
||||
**data,
|
||||
field_id=data["id"],
|
||||
dataset_id=owner,
|
||||
dataset_name=label(dataset) or owner,
|
||||
source="worldquant",
|
||||
collection_version=None,
|
||||
scope=filters.model_dump(include=set(Scope.model_fields)),
|
||||
fetched_at=now(),
|
||||
synced_at=None,
|
||||
)
|
||||
)
|
||||
)
|
||||
count = raw.get("count")
|
||||
known_total = type(count) is int and count >= 0
|
||||
more = (
|
||||
bool(raw["next"])
|
||||
if "next" in raw
|
||||
else (filters.offset + len(items) < count if known_total else len(items) == filters.limit)
|
||||
)
|
||||
result = page(
|
||||
items,
|
||||
count if known_total else filters.offset + len(items) + int(more),
|
||||
filters.limit,
|
||||
filters.offset,
|
||||
)
|
||||
result.update(has_more=more, total_known=known_total)
|
||||
return result
|
||||
|
||||
async def resolve_fields(self, scope, refs):
|
||||
"""Resolve trusted source records before any collection member is written."""
|
||||
if any(ref.scope.key() != scope.key() for ref in refs):
|
||||
raise HTTPException(422, "不能跨区域、Top、Delay 或品种添加字段")
|
||||
result = {}
|
||||
for ref in refs:
|
||||
if ref.source == "local":
|
||||
dataset = await Catalog(self.db).dataset(scope, ref.dataset_id, lock=True)
|
||||
if not dataset.field_version or dataset.field_version != ref.collection_version:
|
||||
raise HTTPException(409, "字段来源已更新,请重新查询后添加")
|
||||
entry = await self.db.get(CatalogEntry, (dataset.field_version, ref.field_id))
|
||||
if not entry:
|
||||
raise HTTPException(422, "字段不属于指定数据集")
|
||||
scope_row = await self.db.get(CatalogScope, scope.key())
|
||||
owner = await self.db.get(CatalogEntry, (scope_row.catalog_version, ref.dataset_id))
|
||||
field = self.local_field(entry, ref.dataset_id, owner.name if owner else None, scope)
|
||||
else:
|
||||
offset, seen, field = 0, set(), None
|
||||
while True:
|
||||
response = await self.online_fields(
|
||||
FieldFilters(
|
||||
**scope.model_dump(),
|
||||
q=ref.field_id,
|
||||
dataset_id=ref.dataset_id,
|
||||
limit=100,
|
||||
offset=offset,
|
||||
)
|
||||
)
|
||||
field = next((item for item in response["items"] if item["id"] == ref.field_id), None)
|
||||
if field or not response["has_more"]:
|
||||
break
|
||||
ids = {item["id"] for item in response["items"]}
|
||||
if not ids - seen:
|
||||
raise HTTPException(502, "平台字段分页未前进")
|
||||
seen.update(ids)
|
||||
offset += len(response["items"])
|
||||
if not field:
|
||||
raise HTTPException(422, "在线字段已不可用,请重新查询")
|
||||
if field["id"] in result and result[field["id"]]["dataset_id"] != field["dataset_id"]:
|
||||
raise HTTPException(422, "同名字段的数据集归属冲突")
|
||||
result[field["id"]] = field
|
||||
return list(result.values())
|
||||
|
||||
async def get(self, preparation_id, version=None, lock=False):
|
||||
query = select(DataPreparation).where(DataPreparation.id == preparation_id)
|
||||
row = await self.db.scalar(query.with_for_update() if lock else query)
|
||||
if not row:
|
||||
raise HTTPException(404, "数据准备集合不存在")
|
||||
if version is not None and row.version != version:
|
||||
raise HTTPException(409, "集合已修改,请重新读取或选择;当前草稿已保留")
|
||||
return row
|
||||
|
||||
async def output(self, row):
|
||||
count, datasets = (
|
||||
await self.db.execute(
|
||||
select(func.count(), func.count(func.distinct(PreparationField.dataset_id))).where(
|
||||
PreparationField.preparation_id == row.id
|
||||
)
|
||||
)
|
||||
).one()
|
||||
return encode_snapshot(
|
||||
dict(
|
||||
id=row.id,
|
||||
name=row.name,
|
||||
note=row.note,
|
||||
scope=row.scope,
|
||||
version=row.version,
|
||||
field_count=count,
|
||||
dataset_count=datasets,
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
)
|
||||
|
||||
async def list(self, q="", scope_key=None, limit=25, offset=0):
|
||||
query = select(DataPreparation)
|
||||
if scope_key:
|
||||
query = query.where(DataPreparation.scope_key == scope_key)
|
||||
if q:
|
||||
query = query.where(
|
||||
or_(
|
||||
DataPreparation.name.ilike(contains(q), escape="\\"),
|
||||
DataPreparation.note.ilike(contains(q), escape="\\"),
|
||||
)
|
||||
)
|
||||
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
|
||||
rows = await self.db.scalars(
|
||||
query.order_by(DataPreparation.updated_at.desc(), DataPreparation.id).limit(limit).offset(offset)
|
||||
)
|
||||
return page([await self.output(row) for row in rows], total, limit, offset)
|
||||
|
||||
async def members(self, preparation_id, q="", dataset_id=None, limit=25, offset=0):
|
||||
await self.get(preparation_id)
|
||||
query = select(PreparationField).where(PreparationField.preparation_id == preparation_id)
|
||||
if dataset_id:
|
||||
query = query.where(PreparationField.dataset_id == dataset_id)
|
||||
if q:
|
||||
query = query.where(
|
||||
or_(
|
||||
PreparationField.field_id.ilike(contains(q), escape="\\"),
|
||||
PreparationField.content["name"].as_string().ilike(contains(q), escape="\\"),
|
||||
PreparationField.content["description"].as_string().ilike(contains(q), escape="\\"),
|
||||
)
|
||||
)
|
||||
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
|
||||
rows = await self.db.scalars(
|
||||
query.order_by(PreparationField.dataset_id, PreparationField.field_id).limit(limit).offset(offset)
|
||||
)
|
||||
return page([row.content for row in rows], total, limit, offset)
|
||||
|
||||
async def create(self, name, note, scope, fields):
|
||||
row = DataPreparation(
|
||||
id=str(uuid4()), name=name, note=note, scope=scope.model_dump(), scope_key=scope.key()
|
||||
)
|
||||
self.db.add(row)
|
||||
await self.db.flush()
|
||||
await self.add(row, fields)
|
||||
return await self.output(row)
|
||||
|
||||
async def add(self, row, fields):
|
||||
for field in fields:
|
||||
existing = await self.db.get(PreparationField, (row.id, field["id"]))
|
||||
if existing:
|
||||
if existing.dataset_id != field["dataset_id"]:
|
||||
raise HTTPException(422, "同名字段的数据集归属冲突")
|
||||
continue
|
||||
self.db.add(
|
||||
PreparationField(
|
||||
preparation_id=row.id, field_id=field["id"], dataset_id=field["dataset_id"], content=field
|
||||
)
|
||||
)
|
||||
await self.db.flush()
|
||||
|
||||
async def copy_dataset(self, body):
|
||||
dataset = await Catalog(self.db).dataset(body.scope, body.dataset_id, lock=True)
|
||||
if not dataset.field_version or dataset.field_version != body.collection_version:
|
||||
raise HTTPException(409, "数据集尚未完整同步或版本已更新")
|
||||
batch = await self.db.get(CatalogBatch, dataset.field_version)
|
||||
if not batch.complete:
|
||||
raise HTTPException(409, "数据集尚未完整同步")
|
||||
source = await Catalog(self.db).detail(body.scope, body.dataset_id)
|
||||
rows = await self.db.scalars(
|
||||
select(CatalogEntry)
|
||||
.where(CatalogEntry.batch_id == dataset.field_version)
|
||||
.order_by(CatalogEntry.id)
|
||||
)
|
||||
fields = [self.local_field(row, body.dataset_id, source["name"], body.scope) for row in rows]
|
||||
return await self.create(
|
||||
f"{source['name'] or body.dataset_id} · {now():%Y%m%d-%H%M%S-%f}", "", body.scope, fields
|
||||
)
|
||||
|
||||
async def remove(self, refs):
|
||||
rows = [await self.get(ref.id, ref.version, lock=True) for ref in sorted(refs, key=lambda r: r.id)]
|
||||
for row in rows:
|
||||
await self.db.execute(delete(PreparationField).where(PreparationField.preparation_id == row.id))
|
||||
await self.db.delete(row)
|
||||
return {"deleted": len(rows)}
|
||||
|
||||
async def freeze(self, refs):
|
||||
"""Lock collection versions and capture source-independent research snapshots atomically."""
|
||||
snapshots = []
|
||||
for ref in sorted(refs, key=lambda r: r.id):
|
||||
row = await self.get(ref.id, ref.version, lock=True)
|
||||
existing = await self.db.scalar(
|
||||
select(ResearchInputSnapshot).where(
|
||||
ResearchInputSnapshot.preparation_id == row.id,
|
||||
ResearchInputSnapshot.preparation_version == row.version,
|
||||
)
|
||||
)
|
||||
if existing:
|
||||
snapshots.append(await self.snapshot(existing.id))
|
||||
continue
|
||||
fields = [
|
||||
r.content
|
||||
for r in await self.db.scalars(
|
||||
select(PreparationField)
|
||||
.where(PreparationField.preparation_id == row.id)
|
||||
.order_by(PreparationField.field_id)
|
||||
)
|
||||
]
|
||||
if not fields:
|
||||
raise HTTPException(422, "空集合不能用于研究")
|
||||
fixed = ResearchInputSnapshot(
|
||||
id=str(uuid4()),
|
||||
preparation_id=row.id,
|
||||
preparation_version=row.version,
|
||||
content=dict(
|
||||
name=row.name,
|
||||
scope=row.scope,
|
||||
fields=fields,
|
||||
field_ids=[f["id"] for f in fields],
|
||||
field_types={f["id"]: f["field_type"] for f in fields},
|
||||
dataset_ids=sorted({f["dataset_id"] for f in fields}),
|
||||
),
|
||||
)
|
||||
self.db.add(fixed)
|
||||
await self.db.flush()
|
||||
snapshots.append(await self.snapshot(fixed.id))
|
||||
return snapshots
|
||||
|
||||
async def snapshot(self, snapshot_id):
|
||||
row = await self.db.get(ResearchInputSnapshot, snapshot_id)
|
||||
if not row:
|
||||
raise HTTPException(404, "研究输入快照不存在")
|
||||
return encode_snapshot(
|
||||
dict(
|
||||
**row.content,
|
||||
id=row.id,
|
||||
preparation_id=row.preparation_id,
|
||||
preparation_version=row.preparation_version,
|
||||
created_at=row.created_at,
|
||||
)
|
||||
)
|
||||
|
||||
async def bind(self, body):
|
||||
refs = getattr(body, "preparation_refs", [])
|
||||
if refs:
|
||||
fixed = await self.freeze(refs)
|
||||
body.input_ids = list(dict.fromkeys([*body.input_ids, *[r["id"] for r in fixed]]))
|
||||
body.preparation_refs = []
|
||||
if not body.input_ids:
|
||||
raise HTTPException(422, "请选择非空的数据准备集合")
|
||||
limit = next(
|
||||
m.max_length for m in type(body).model_fields["input_ids"].metadata if hasattr(m, "max_length")
|
||||
)
|
||||
if len(body.input_ids) > limit:
|
||||
raise HTTPException(422, f"最多可选择 {limit} 个研究输入")
|
||||
return body
|
||||
@@ -4,6 +4,8 @@ from pydantic import Field
|
||||
|
||||
from ..ai.alpha_tools import AlphaArgs
|
||||
from ..ai.capabilities import Capability
|
||||
from ..preparations.service import Preparations
|
||||
from ..schemas import Contract
|
||||
from .contracts import ChatboxResearchInput, InputPageArgs, ResearchInputSelection, ResearchPreviewInput
|
||||
|
||||
|
||||
@@ -16,14 +18,26 @@ async def prepare(ctx, args):
|
||||
return await ctx.business.research_builder.prepare(ResearchPreviewInput(**args.model_dump()))
|
||||
|
||||
|
||||
INSTRUCTIONS = "Chatbox 是研究来源,不是 Alpha 备注。新候选由服务端记录会话和生成轮次;引用已有草稿和重跑保留原来源。\n数据集研究先用 search_catalog 读取实际字段/类型/集合版本,或 get_research_input 读取页面提供的固定输入。\n只有明确选出的字段才能 prepare_research_input;不要把一页搜索结果当作整集。页面 unsaved_field_selection 为 true 且无输入引用时,请用户保存选择并点击“用此输入研究”,不能忽略排除项。\n有输入快照时使用 prepare_research_backtest,提供研究假设、具名占位符和实际字段类型,模拟参数须匹配输入范围。模板中的所有数据字段使用绑定占位符。\n字段类型未知时说明限制;VECTOR 处理方式在模板中明确写出,不把 VECTOR 当作 MATRIX。绑定校验不等于 FASTEXPR 语义或平台权限验证。\n无数据集输入的直接表达式研究可以 prepare_backtest;不得声称来自已核验数据集。已有候选草稿先 get_backtest_draft,再按 ID 和版本引用。"
|
||||
INSTRUCTIONS = "研究先用 search_data_preparations 查询可编辑集合,读取集合 ID 与 version 后使用 prepare_research_input 固定输入。已有快照使用 get_research_input。所有研究来源保留独立快照,删除集合不影响已有研究。构建回测需明确假设、字段绑定和范围;VECTOR 必须显式处理,不能当作 MATRIX。直接表达式回测不声明数据准备来源。"
|
||||
|
||||
|
||||
|
||||
class PreparationSearch(Contract):
|
||||
q: str = Field(default="", max_length=300)
|
||||
scope_key: str | None = None
|
||||
limit: int = Field(default=25, ge=1, le=100)
|
||||
offset: int = Field(default=0, ge=0)
|
||||
|
||||
|
||||
CAPABILITIES = (
|
||||
Capability(name="search_data_preparations", schema=PreparationSearch,
|
||||
description="分页搜索数据准备集合,返回 ID、version、范围与字段数;非空集合可固定为研究输入。",
|
||||
label="查询数据准备", renderer="catalog", effect="query",
|
||||
handler=lambda ctx, args: Preparations(ctx.business.db).list(**args.model_dump())),
|
||||
Capability(
|
||||
name="prepare_research_input",
|
||||
schema=ResearchInputSelection,
|
||||
description="把明确选出的 1–100 个字段固定为研究输入快照;必须提供已读取的集合版本。只保存本地输入,不启动同步或回测。已有输入引用时直接读取,不重建。",
|
||||
description="将数据准备集合的明确 ID 和 version 固定为研究快照,保留完整字段与数据集归属。已有快照直接读取。",
|
||||
label="固定研究输入",
|
||||
renderer="catalog",
|
||||
effect="prepare",
|
||||
|
||||
@@ -57,7 +57,11 @@ class Assets:
|
||||
"view": ViewSpec,
|
||||
"workflow": WorkflowSpec,
|
||||
}[body.kind]
|
||||
content = schema.model_validate(body.content).model_dump(mode="json")
|
||||
parsed = schema.model_validate(body.content)
|
||||
if body.kind == "feature":
|
||||
from ..preparations.service import Preparations
|
||||
await Preparations(self.db).bind(parsed)
|
||||
content = parsed.model_dump(mode="json")
|
||||
if body.kind == "workflow":
|
||||
from .workflows import validate_graph
|
||||
|
||||
|
||||
@@ -5,16 +5,13 @@ from typing import Literal
|
||||
from pydantic import Field, model_validator
|
||||
|
||||
from ..backtests.contracts import SimulationSettings, Source
|
||||
from ..catalog.contracts import Scope
|
||||
from ..preparations.contracts import PreparationReference
|
||||
from ..schemas import Contract
|
||||
from .expressions import PLACEHOLDER
|
||||
|
||||
|
||||
class ResearchInputSelection(Contract):
|
||||
scope: Scope
|
||||
dataset_id: str = Field(min_length=1, max_length=200)
|
||||
collection_version: str = Field(min_length=1, max_length=36)
|
||||
field_ids: list[str] = Field(min_length=1, max_length=100)
|
||||
items: list[PreparationReference] = Field(min_length=1, max_length=1)
|
||||
|
||||
|
||||
class InputPageArgs(Contract):
|
||||
@@ -50,7 +47,7 @@ class ChatboxResearchInput(Contract):
|
||||
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
hypothesis: str = Field(min_length=1, max_length=2000)
|
||||
template_input_id: str = Field(min_length=1, max_length=36)
|
||||
input_snapshot_id: str = Field(min_length=1, max_length=36)
|
||||
candidates: list[ResearchCandidate] = Field(min_length=1, max_length=100)
|
||||
|
||||
@model_validator(mode="after")
|
||||
|
||||
@@ -11,7 +11,8 @@ from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Simulatio
|
||||
from ..backtests.service import Backtests, uid
|
||||
from ..catalog.research_metadata import ResearchMetadata
|
||||
from ..catalog.service import Catalog
|
||||
from ..models import Alpha, BacktestRun, CatalogResource, ResearchExperiment, TemplateInput
|
||||
from ..models import Alpha, BacktestRun, CatalogResource, ResearchExperiment
|
||||
from ..preparations.service import Preparations
|
||||
from .assets import Assets
|
||||
from .expressions import GROUPS, analyze, expand
|
||||
from .serialization import encode_snapshot as jsonable_encoder
|
||||
@@ -86,7 +87,7 @@ class Experiments:
|
||||
"candidates": experiment["candidates"],
|
||||
"hypothesis": experiment["hypothesis"],
|
||||
"input_references": [
|
||||
{k: entry[k] for k in ("id", "collection_version", "scope", "dataset_id")}
|
||||
{k: entry[k] for k in ("id", "preparation_id", "preparation_version", "scope", "dataset_ids")}
|
||||
for entry in experiment["inputs"]
|
||||
],
|
||||
"template_reference": {
|
||||
@@ -133,6 +134,7 @@ class Experiments:
|
||||
return validation
|
||||
|
||||
async def create(self, body, kind="template", extra_evidence=None, *, parent_snapshots=None):
|
||||
await Preparations(self.db).bind(body)
|
||||
asset = await self.assets.get(body.asset_id, body.version, "template") if body.asset_id else None
|
||||
template = TemplateSpec.model_validate(asset["content"]) if asset else body.template
|
||||
scope = scope_of(body.settings)
|
||||
@@ -318,7 +320,8 @@ class Experiments:
|
||||
kind=source_kind or experiment["kind"],
|
||||
reference=reference or experiment_id,
|
||||
research_id=experiment_id,
|
||||
template_input_id=inputs[0]["id"] if len(inputs) == 1 else None,
|
||||
input_snapshot_ids=[i["id"] for i in inputs],
|
||||
input_snapshot_id=inputs[0]["id"] if len(inputs) == 1 else None,
|
||||
hypothesis=experiment["hypothesis"][:2000],
|
||||
),
|
||||
candidates=[
|
||||
@@ -342,6 +345,7 @@ class Experiments:
|
||||
original = parents[0]
|
||||
base = seed_settings(original["settings"])
|
||||
expression = original["expression"]
|
||||
await Preparations(self.db).bind(body)
|
||||
snapshots, _ = await self.inputs(body.input_ids)
|
||||
groups = defaultdict(list)
|
||||
for snapshot in snapshots:
|
||||
@@ -405,6 +409,7 @@ class Experiments:
|
||||
)
|
||||
|
||||
async def generation_context(self, body):
|
||||
await Preparations(self.db).bind(body)
|
||||
snapshots, fields = await self.inputs(body.input_ids)
|
||||
parents = await self.parents(body.parent_alpha_ids, body.parent_experiment_ids)
|
||||
metadata = await ResearchMetadata(self.db).operators(limit=100)
|
||||
@@ -414,7 +419,7 @@ class Experiments:
|
||||
"hypothesis": body.hypothesis,
|
||||
"method": body.method,
|
||||
"inputs": [
|
||||
{"id": item["id"], "scope": item["scope"], "dataset_id": item["dataset_id"]}
|
||||
{"id": item["id"], "scope": item["scope"], "dataset_ids": item["dataset_ids"], "name": item["name"], "fields": item["fields"][:100]}
|
||||
for item in snapshots
|
||||
],
|
||||
"fields": dict(list(fields.items())[:300]),
|
||||
@@ -424,9 +429,3 @@ class Experiments:
|
||||
],
|
||||
"parents": [{**p, "candidates": p.get("candidates", [])[:10]} for p in parents],
|
||||
}
|
||||
|
||||
async def available_inputs(self, limit=100):
|
||||
rows = await self.db.scalars(
|
||||
select(TemplateInput).order_by(TemplateInput.created_at.desc()).limit(limit)
|
||||
)
|
||||
return {"items": [await self.catalog.input(row.id) for row in rows]}
|
||||
|
||||
@@ -28,12 +28,6 @@ from .workspace_contracts import (
|
||||
router = APIRouter(prefix="/api/v1/research", tags=["research"], dependencies=[Depends(require_auth)])
|
||||
|
||||
|
||||
@router.get("/inputs")
|
||||
async def inputs(request: Request, limit: int = Query(100, ge=1, le=100)):
|
||||
async with request.app.state.sessions() as db:
|
||||
return await Experiments(db).available_inputs(limit)
|
||||
|
||||
|
||||
@router.get("/assets")
|
||||
async def assets(
|
||||
request: Request,
|
||||
@@ -96,10 +90,10 @@ async def import_commit(body: ImportCommit, request: Request):
|
||||
|
||||
@router.post("/generate", status_code=201)
|
||||
async def generate(body: Generation, request: Request):
|
||||
async with request.app.state.sessions() as db:
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
context = await Experiments(db).generation_context(body)
|
||||
result, evidence = await request_model(request.app.state.ai, context, OUTPUTS[body.method])
|
||||
if body.method == "feature" and set(result.input_ids) != set(body.input_ids):
|
||||
if body.method == "feature" and (result.preparation_refs or set(result.input_ids) != set(body.input_ids)):
|
||||
raise HTTPException(422, "模型不能改变已固定的输入范围")
|
||||
async with request.app.state.sessions.begin() as db:
|
||||
asset = await Assets(db).save(
|
||||
|
||||
@@ -476,9 +476,9 @@ class ResearchRuntime:
|
||||
step = await db.get(ResearchStepRun, step_id)
|
||||
if not step or step.status != "running":
|
||||
return
|
||||
if isinstance(result, FeatureSpec) and set(result.input_ids) != set(
|
||||
if isinstance(result, FeatureSpec) and (result.preparation_refs or set(result.input_ids) != set(
|
||||
[i["id"] for i in step.output["context"]["inputs"]]
|
||||
):
|
||||
)):
|
||||
raise HTTPException(422, "模型不能改变已固定的输入范围")
|
||||
# A paused/stopped run may collect this already-issued model output, but cannot advance.
|
||||
asset = await Assets(db).save(
|
||||
|
||||
@@ -5,12 +5,9 @@ not FASTEXPR operator semantics or the account's current platform permissions.
|
||||
"""
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
|
||||
from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Source
|
||||
from ..catalog.contracts import EntryOutput, InputPreparation
|
||||
from ..catalog.service import Catalog
|
||||
from ..models import CatalogEntry
|
||||
from .expressions import analyze, expand
|
||||
|
||||
|
||||
@@ -21,55 +18,18 @@ class ResearchBuilder:
|
||||
self.backtests = backtests
|
||||
|
||||
async def select_input(self, body):
|
||||
"""Fix explicit fields in one published version; reject missing or stale members."""
|
||||
collection = await self.catalog.collection(body.scope, body.dataset_id)
|
||||
chosen = set(body.field_ids)
|
||||
if len(chosen) != len(body.field_ids) or not chosen.issubset(collection["field_ids"]):
|
||||
raise HTTPException(422, "字段选择含重复、未知或其他数据集字段")
|
||||
saved = await self.catalog.prepare(
|
||||
InputPreparation(
|
||||
scope=body.scope,
|
||||
dataset_id=body.dataset_id,
|
||||
collection_version=body.collection_version,
|
||||
selection="explicit",
|
||||
excluded_ids=[field for field in collection["field_ids"] if field not in chosen],
|
||||
)
|
||||
)
|
||||
from ..preparations.service import Preparations
|
||||
saved = (await Preparations(self.db).freeze(body.items))[0]
|
||||
return await self.input_page(saved["id"])
|
||||
|
||||
async def input_page(self, input_id, limit=25, offset=0, q="", field_type=None):
|
||||
"""Read the saved version, including field descriptions, with explicit pagination."""
|
||||
saved = await self.catalog.input(input_id)
|
||||
ids = [
|
||||
field
|
||||
for field in saved["field_ids"]
|
||||
if q.lower() in field.lower()
|
||||
and (field_type is None or saved["field_types"].get(field) == field_type)
|
||||
]
|
||||
page = ids[offset : offset + limit]
|
||||
entries = {
|
||||
row.id: row
|
||||
for row in await self.db.scalars(
|
||||
select(CatalogEntry).where(
|
||||
CatalogEntry.batch_id == saved["collection_version"], CatalogEntry.id.in_(page)
|
||||
)
|
||||
)
|
||||
}
|
||||
return {
|
||||
**{
|
||||
k: saved[k]
|
||||
for k in ("id", "scope", "dataset_id", "collection_version", "selection", "created_at")
|
||||
},
|
||||
"field_count": len(saved["field_ids"]),
|
||||
"items": [
|
||||
EntryOutput.model_validate(entries[field], from_attributes=True).model_dump()
|
||||
for field in page
|
||||
],
|
||||
"total": len(ids),
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
"has_more": offset + limit < len(ids),
|
||||
}
|
||||
fields = [f for f in saved["fields"] if (not q or q.lower() in
|
||||
" ".join(str(f.get(k) or "") for k in ("id", "name", "description", "dataset_id")).lower())
|
||||
and (not field_type or f["field_type"] == field_type)]
|
||||
return {**{k: v for k, v in saved.items() if k not in ("fields", "field_ids", "field_types")},
|
||||
"field_count": len(saved["fields"]), "items": fields[offset:offset + limit],
|
||||
"total": len(fields), "limit": limit, "offset": offset, "has_more": offset + limit < len(fields)}
|
||||
|
||||
async def prepare(self, body):
|
||||
"""Bind templates against an immutable input, then reuse the fixed-preview interface.
|
||||
@@ -77,7 +37,7 @@ class ResearchBuilder:
|
||||
Raises HTTPException(422) for wrong scope, membership or declared type.
|
||||
No expression execution or implicit cleaning/aggregation takes place here.
|
||||
"""
|
||||
saved = await self.catalog.input(body.template_input_id)
|
||||
saved = await self.catalog.input(body.input_snapshot_id)
|
||||
scope = saved["scope"]
|
||||
candidates = []
|
||||
for item in body.candidates:
|
||||
@@ -114,7 +74,7 @@ class ResearchBuilder:
|
||||
source = Source.model_validate(
|
||||
{
|
||||
**body.source.model_dump(),
|
||||
"template_input_id": saved["id"],
|
||||
"input_snapshot_id": saved["id"],
|
||||
"hypothesis": body.hypothesis,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -181,6 +181,8 @@ class Workflows:
|
||||
for node in graph.nodes:
|
||||
if node.type == "iterate" and node.config["max_rounds"] > body.budget.max_rounds:
|
||||
raise HTTPException(422, "流程迭代上限超过本次授权轮数")
|
||||
from ..preparations.service import Preparations
|
||||
await Preparations(self.db).bind(body)
|
||||
experiments = Experiments(self.db)
|
||||
settings_variant = any(
|
||||
n.type == "variant" and n.config.get("method") == "settings" for n in graph.nodes
|
||||
|
||||
@@ -7,6 +7,7 @@ from pydantic import Field, field_validator, model_validator
|
||||
|
||||
from ..backtests.contracts import SimulationSettings
|
||||
from ..catalog.contracts import Scope
|
||||
from ..preparations.contracts import PreparationReference
|
||||
from ..schemas import Contract
|
||||
from .expressions import IDENTIFIER, PLACEHOLDER, normalize_template
|
||||
|
||||
@@ -70,7 +71,8 @@ class FeatureStep(Contract):
|
||||
class FeatureSpec(Contract):
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
hypothesis: str = Field(min_length=1, max_length=10000)
|
||||
input_ids: list[str] = Field(min_length=1, max_length=20)
|
||||
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
||||
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
||||
steps: list[FeatureStep] = Field(default_factory=list, max_length=30)
|
||||
template: TemplateSpec | None = None
|
||||
|
||||
@@ -102,7 +104,8 @@ class Expansion(Contract):
|
||||
asset_id: str | None = Field(default=None, max_length=36)
|
||||
version: int | None = Field(default=None, ge=1)
|
||||
template: TemplateSpec | None = None
|
||||
input_ids: list[str] = Field(min_length=1, max_length=20)
|
||||
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
||||
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
||||
hypothesis: str = Field(min_length=1, max_length=10000)
|
||||
settings: SimulationSettings
|
||||
mode: Literal["all", "random"] = "all"
|
||||
@@ -123,7 +126,8 @@ class Expansion(Contract):
|
||||
class Generation(Contract):
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
hypothesis: str = Field(min_length=1, max_length=10000)
|
||||
input_ids: list[str] = Field(min_length=1, max_length=20)
|
||||
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
||||
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
||||
parent_alpha_ids: list[str] = Field(default_factory=list, max_length=20)
|
||||
parent_experiment_ids: list[str] = Field(default_factory=list, max_length=20)
|
||||
method: Literal["template", "structure", "feature"] = "template"
|
||||
@@ -131,7 +135,8 @@ class Generation(Contract):
|
||||
|
||||
class SettingVariants(Contract):
|
||||
alpha_id: str = Field(min_length=1, max_length=100)
|
||||
input_ids: list[str] = Field(min_length=1, max_length=100)
|
||||
input_ids: list[str] = Field(default_factory=list, max_length=100)
|
||||
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=100)
|
||||
hypothesis: str = Field(default="保持表达式,比较字段共同支持的市场与设置", max_length=10000)
|
||||
|
||||
|
||||
@@ -229,7 +234,8 @@ class FlowStart(Contract):
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
workflow_id: str | None = None
|
||||
workflow_version: int | None = Field(default=None, ge=1)
|
||||
input_ids: list[str] = Field(min_length=1, max_length=20)
|
||||
input_ids: list[str] = Field(default_factory=list, max_length=20)
|
||||
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
||||
hypothesis: str = Field(min_length=1, max_length=10000)
|
||||
settings: SimulationSettings
|
||||
budget: Budget
|
||||
|
||||
@@ -7,6 +7,7 @@ from pydantic import Field, model_validator
|
||||
|
||||
from ..backtests.contracts import Candidate, SimulationSettings
|
||||
from ..catalog.contracts import CatalogFilters, Scope
|
||||
from ..preparations.contracts import PreparationReference
|
||||
from ..research.workspace_contracts import TemplateSpec
|
||||
from ..schemas import Contract
|
||||
|
||||
@@ -54,8 +55,24 @@ class Provenance(Contract):
|
||||
parent_run_id: RunId | None = None
|
||||
|
||||
|
||||
class PreparationSearch(Contract):
|
||||
q: str = Field(default="", max_length=300)
|
||||
scope_key: str | None = None
|
||||
limit: int = Field(default=25, ge=1, le=100)
|
||||
offset: int = Field(default=0, ge=0)
|
||||
|
||||
|
||||
class PreparationRead(Contract):
|
||||
id: str = Field(min_length=1, max_length=36)
|
||||
version: int = Field(ge=1)
|
||||
q: str = Field(default="", max_length=300)
|
||||
limit: int = Field(default=25, ge=1, le=100)
|
||||
offset: int = Field(default=0, ge=0)
|
||||
|
||||
|
||||
class Submit(Contract):
|
||||
name: str = Field(min_length=1, max_length=200)
|
||||
preparation_refs: list[PreparationReference] = Field(default_factory=list, max_length=20)
|
||||
candidates: list[DirectCandidate] = Field(min_length=1, max_length=100)
|
||||
idempotency_key: Identifier
|
||||
duplicate_policy: Literal["reject", "rerun"] = "reject"
|
||||
|
||||
@@ -118,6 +118,17 @@ class ResearchAccess:
|
||||
return {"job_id": job.id, "status": job.status, "action": job.kind,
|
||||
"read_with": "get_worldquant_connection", "web_url": f"{self.public_origin}/"}
|
||||
|
||||
async def preparations(self, args):
|
||||
from ..preparations.service import Preparations
|
||||
return await Preparations(self.db).list(args.q, args.scope_key, args.limit, args.offset)
|
||||
|
||||
async def preparation(self, args):
|
||||
from ..preparations.service import Preparations
|
||||
service = Preparations(self.db)
|
||||
row = await service.get(args.id, args.version, lock=True)
|
||||
return {"collection": await service.output(row),
|
||||
"fields": await service.members(row.id, args.q, None, args.limit, args.offset)}
|
||||
|
||||
async def catalog(self, args):
|
||||
data = await Catalog(self.db).search(args.filters, args.dataset_id)
|
||||
return {**data, "scope": args.filters.model_dump(include={"region", "universe", "delay", "instrument_type"}),
|
||||
@@ -298,7 +309,8 @@ class ResearchAccess:
|
||||
# preserve_source prevents the Chatbox-specific generating context rewriting MCP provenance.
|
||||
backtests = Backtests(self.db, provenance)
|
||||
preview = await backtests.preview(PreviewInput(inline=DraftInput(
|
||||
name=args.name, source=source, candidates=args.candidates)), preserve_source=True)
|
||||
name=args.name, source=source, candidates=args.candidates,
|
||||
preparation_refs=args.preparation_refs)), preserve_source=True)
|
||||
result = await backtests.start(StartInput(preview_id=preview["preview_id"],
|
||||
idempotency_key="mcp-" + str(uuid4())))
|
||||
result = {**result, "input_digest": digest, "batch_count": preview["batch_count"],
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
"""Replace saved dataset inputs with editable preparations and independent snapshots."""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "0015"
|
||||
down_revision = "0014"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.drop_table("template_inputs")
|
||||
op.add_column("catalog_batches", sa.Column("job_id", sa.String(36), nullable=True))
|
||||
op.add_column("catalog_batches", sa.Column("offset", sa.Integer(), nullable=False, server_default="0"))
|
||||
op.execute("UPDATE catalog_batches SET job_id = id")
|
||||
# Preserve unfinished catalog pagination; only legacy research inputs are discarded.
|
||||
jobs = sa.table("sync_jobs", sa.column("id"), sa.column("checkpoint", sa.JSON()))
|
||||
batches = sa.table("catalog_batches", sa.column("id"), sa.column("offset", sa.Integer()))
|
||||
for job_id, checkpoint in op.get_bind().execute(sa.select(jobs.c.id, jobs.c.checkpoint)):
|
||||
offset = (checkpoint or {}).get("offset", 0)
|
||||
if type(offset) is int and offset >= 0:
|
||||
op.get_bind().execute(batches.update().where(batches.c.id == job_id).values(offset=offset))
|
||||
names = {"fk": "fk_%(table_name)s_%(column_0_name)s_%(referred_table_name)s"}
|
||||
fk = next(
|
||||
f
|
||||
for f in sa.inspect(op.get_bind()).get_foreign_keys("catalog_batches")
|
||||
if f["constrained_columns"] == ["id"]
|
||||
)
|
||||
with op.batch_alter_table("catalog_batches", naming_convention=names) as batch:
|
||||
batch.drop_constraint(fk["name"] or "fk_catalog_batches_id_sync_jobs", type_="foreignkey")
|
||||
batch.alter_column("job_id", existing_type=sa.String(36), nullable=False)
|
||||
batch.create_foreign_key("fk_catalog_batches_job", "sync_jobs", ["job_id"], ["id"])
|
||||
batch.create_index("ix_catalog_batches_job_id", ["job_id"])
|
||||
op.create_table(
|
||||
"data_preparations",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("note", sa.Text(), nullable=False),
|
||||
sa.Column("scope_key", sa.String(200), nullable=False),
|
||||
sa.Column("scope", sa.JSON(), nullable=False),
|
||||
sa.Column("version", sa.Integer(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
)
|
||||
op.create_index("ix_data_preparations_scope_key", "data_preparations", ["scope_key"])
|
||||
op.create_table(
|
||||
"preparation_fields",
|
||||
sa.Column(
|
||||
"preparation_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("data_preparations.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
),
|
||||
sa.Column("field_id", sa.String(200), primary_key=True),
|
||||
sa.Column("dataset_id", sa.String(200), nullable=False),
|
||||
sa.Column("content", sa.JSON(), nullable=False),
|
||||
)
|
||||
op.create_index("ix_preparation_fields_dataset_id", "preparation_fields", ["dataset_id"])
|
||||
op.create_table(
|
||||
"research_input_snapshots",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("preparation_id", sa.String(36), nullable=False),
|
||||
sa.Column("preparation_version", sa.Integer(), nullable=False),
|
||||
sa.Column("content", sa.JSON(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.UniqueConstraint("preparation_id", "preparation_version"),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_research_input_snapshots_preparation_id", "research_input_snapshots", ["preparation_id"]
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
raise RuntimeError("旧输入模型已移除;回退请恢复升级前数据库备份")
|
||||
@@ -43,8 +43,7 @@ if __name__ == "__main__":
|
||||
"INSERT INTO research (alpha_id, note, tags, favorite, state, updated_at, version) VALUES ('MIGRATION_TEST', 'preserve research', '[]', false, 'inbox', now(), 7);"
|
||||
)
|
||||
)
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
command.upgrade(config, "0014")
|
||||
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
|
||||
("preserve research", 7)
|
||||
]
|
||||
@@ -56,7 +55,7 @@ if __name__ == "__main__":
|
||||
("preserve research", 7)
|
||||
]
|
||||
print(
|
||||
"PostgreSQL 17: 0002 → 0003, downgrade/re-upgrade, metadata check, Alpha/research preservation passed"
|
||||
"PostgreSQL 17: catalog migration, 0014 downgrade/re-upgrade then head, metadata and Alpha preservation passed"
|
||||
)
|
||||
|
||||
async def flow():
|
||||
@@ -73,7 +72,6 @@ if __name__ == "__main__":
|
||||
return httpx.Response(201, json={"token": {"expiry": 14400}})
|
||||
if request.url.path == "/users/self":
|
||||
return httpx.Response(200, json={"id": "PG_TEST_USER"})
|
||||
assert request.method == "GET"
|
||||
return catalog_response(request) or httpx.Response(404)
|
||||
|
||||
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
|
||||
@@ -115,7 +113,7 @@ if __name__ == "__main__":
|
||||
assert sorted(r.status_code for r in responses) == [200, 409]
|
||||
await sync(catalog, "TEST_FIN")
|
||||
assert (await prepare(client, version)).status_code == 409
|
||||
persisted = (await client.get("/api/v1/catalog/inputs/" + draft["id"])).json()
|
||||
persisted = (await client.get("/api/v1/research/input-snapshots/" + draft["id"])).json()
|
||||
assert persisted == draft
|
||||
print(
|
||||
"PostgreSQL: real API/runner multi-page dedupe, immutable draft, refresh conflict and concurrent note CAS passed"
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
"""Isolated PostgreSQL migration and concurrency acceptance, using synthetic upstream only."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from cryptography.fernet import Fernet
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
URL = "postgresql+asyncpg://postgres:preparations-test-only@127.0.0.1:18437/preparations_test"
|
||||
os.environ.update(
|
||||
DATABASE_URL=URL,
|
||||
ADMIN_PASSWORD="migration-test-only",
|
||||
ENCRYPTION_KEY=Fernet.generate_key().decode(),
|
||||
WQ_EMAIL="",
|
||||
WQ_PASSWORD="",
|
||||
)
|
||||
|
||||
|
||||
async def sql(query):
|
||||
engine = create_async_engine(URL)
|
||||
try:
|
||||
async with engine.begin() as db:
|
||||
result = await db.execute(text(query))
|
||||
return result.fetchall() if result.returns_rows else None
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def acceptance():
|
||||
import httpx
|
||||
|
||||
from app.catalog.contracts import CatalogJobInput, Scope
|
||||
from app.catalog.service import Catalog
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.worldquant import WqClient
|
||||
from tests.catalog_fake import catalog_response
|
||||
from tests.test_catalog import SCOPE, prepare, search, sync
|
||||
|
||||
def upstream(request):
|
||||
if request.url.path == "/authentication":
|
||||
return httpx.Response(201, json={"token": {"expiry": 14400}})
|
||||
if request.url.path == "/users/self":
|
||||
return httpx.Response(200, json={"id": "PG_TEST_USER"})
|
||||
from tests.research_metadata_fake import response
|
||||
|
||||
metadata = response(request)
|
||||
if metadata is not None:
|
||||
return metadata
|
||||
assert request.method == "GET"
|
||||
return catalog_response(request) or httpx.Response(404)
|
||||
|
||||
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
|
||||
app = create_app(settings, WqClient(settings, transport=httpx.MockTransport(upstream)))
|
||||
async with app.router.lifespan_context(app):
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=app),
|
||||
base_url="http://testserver",
|
||||
headers={"X-WQ-Request": "1"},
|
||||
) as client:
|
||||
assert (
|
||||
await client.post(
|
||||
"/api/v1/auth/login", json={"username": "admin", "password": "migration-test-only"}
|
||||
)
|
||||
).status_code == 200
|
||||
await client.put(
|
||||
"/api/v1/account/credentials",
|
||||
json={"email": "synthetic@example.com", "password": "synthetic-only"},
|
||||
)
|
||||
job = (await client.post("/api/v1/account/connect")).json()
|
||||
await app.state.runner.execute(job["id"])
|
||||
fixture = (client, app.state.runner, {})
|
||||
await sync(fixture)
|
||||
await sync(fixture, "TEST_FIN")
|
||||
snapshot = (
|
||||
await prepare(
|
||||
client, (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"]
|
||||
)
|
||||
).json()
|
||||
assert len(snapshot["field_ids"]) == 123
|
||||
fields = await client.get("/api/v1/catalog/fields", params={**SCOPE, "category": "基本面", "limit": 2, "offset": 2})
|
||||
assert fields.status_code == 200, fields.text
|
||||
assert fields.json()["total"] == 123 and len(fields.json()["items"]) == 2
|
||||
assert fields.json()["items"][0]["category"] == "基本面"
|
||||
|
||||
|
||||
async def enqueue():
|
||||
async with app.state.sessions.begin() as db:
|
||||
return (await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)).id
|
||||
|
||||
jobs = await asyncio.gather(*(enqueue() for _ in range(5)))
|
||||
assert len(set(jobs)) == 1, jobs
|
||||
collection = (await client.get("/api/v1/data-preparations")).json()["items"][0]
|
||||
ref = {"id": collection["id"], "version": collection["version"]}
|
||||
|
||||
async def freeze():
|
||||
response = await client.post("/api/v1/data-preparations/freeze", json={"items": [ref]})
|
||||
assert response.status_code == 201, response.text
|
||||
return response.json()["items"][0]["id"]
|
||||
|
||||
assert len(set(await asyncio.gather(*(freeze() for _ in range(5))))) == 1
|
||||
await client.delete(f"/api/v1/data-preparations/{ref['id']}?version={ref['version']}")
|
||||
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
assert not asyncio.run(sql("SELECT tablename FROM pg_tables WHERE schemaname='public'")), (
|
||||
"Requires an empty isolated database"
|
||||
)
|
||||
config = Config("alembic.ini")
|
||||
command.upgrade(config, "0014")
|
||||
# Existing catalog checkpoints survive the structural change; no research data migration.
|
||||
asyncio.run(
|
||||
sql(
|
||||
"INSERT INTO sync_jobs(id,kind,status,payload,checkpoint,total,processed,failed,cancel_requested,created_at,updated_at) VALUES ('checkpoint-test','catalog_sync','failed','{}','{\"offset\": 100}',0,100,0,false,now(),now())"
|
||||
)
|
||||
)
|
||||
asyncio.run(sql("INSERT INTO catalog_scopes(key,scope) VALUES ('checkpoint-scope','{}')"))
|
||||
asyncio.run(
|
||||
sql(
|
||||
"INSERT INTO catalog_batches(id,scope_key,dataset_id,complete,count) VALUES ('checkpoint-test','checkpoint-scope',NULL,false,100)"
|
||||
)
|
||||
)
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
assert asyncio.run(
|
||||
sql("SELECT job_id, catalog_batches.\"offset\" FROM catalog_batches WHERE id='checkpoint-test'")
|
||||
) == [("checkpoint-test", 100)]
|
||||
assert asyncio.run(sql("SELECT to_regclass('template_inputs')")) == [(None,)]
|
||||
asyncio.run(sql("DELETE FROM catalog_batches WHERE id='checkpoint-test'"))
|
||||
asyncio.run(sql("DELETE FROM catalog_scopes WHERE key='checkpoint-scope'"))
|
||||
asyncio.run(sql("DELETE FROM sync_jobs WHERE id='checkpoint-test'"))
|
||||
asyncio.run(acceptance())
|
||||
print(
|
||||
"PostgreSQL 17: 0014 → 0015 metadata, retained catalog checkpoint, concurrent job deduplication/freeze and independent snapshot passed"
|
||||
)
|
||||
@@ -22,7 +22,7 @@ def research_step(text, returns, history):
|
||||
return "get_backtest_results", {"run_id": run_id}
|
||||
data = content(returns[-1])
|
||||
return f"已读取本次研究的 {data['total']} 条真实保存结果;缺失 Sharpe 仍为未知。"
|
||||
if context.get("unsaved_field_selection") and not context.get("template_input_id"):
|
||||
if context.get("unsaved_field_selection") and not context.get("input_snapshot_id"):
|
||||
return "请先保存字段选择,再点击用此输入研究。"
|
||||
if not returns:
|
||||
return "get_backtest_capabilities", {}
|
||||
@@ -30,39 +30,23 @@ def research_step(text, returns, history):
|
||||
data = content(last)
|
||||
if "error" in data:
|
||||
return f"研究尚未完成:{data['error']}"
|
||||
scope = context.get("catalog_scope") or {
|
||||
"instrument_type": "EQUITY",
|
||||
"region": "USA",
|
||||
"universe": "TOP3000",
|
||||
"delay": 1,
|
||||
}
|
||||
if last.tool_name == "get_backtest_capabilities":
|
||||
if context.get("template_input_id"):
|
||||
if context.get("input_snapshot_id"):
|
||||
return "get_research_input", {
|
||||
"input_id": context["template_input_id"],
|
||||
"input_id": context["input_snapshot_id"],
|
||||
"field_type": "MATRIX",
|
||||
"limit": 1,
|
||||
}
|
||||
return "search_catalog", {"filters": {**scope, "q": "TEST_FIN", "limit": 1}}
|
||||
if last.tool_name == "search_catalog":
|
||||
if data["dataset_id"] is None:
|
||||
return "search_catalog", {
|
||||
"dataset_id": data["items"][0]["id"],
|
||||
"filters": {**scope, "field_type": "MATRIX", "limit": 1},
|
||||
}
|
||||
return "prepare_research_input", {
|
||||
"scope": scope,
|
||||
"dataset_id": data["dataset_id"],
|
||||
"collection_version": data["collection_version"],
|
||||
"field_ids": [data["items"][0]["id"]],
|
||||
}
|
||||
return "search_data_preparations", {"limit": 1}
|
||||
if last.tool_name == "search_data_preparations":
|
||||
return "prepare_research_input", {"items": [{"id": data["items"][0]["id"], "version": data["items"][0]["version"]}]}
|
||||
if last.tool_name in ("get_research_input", "prepare_research_input"):
|
||||
field = data["items"][0]
|
||||
field = next(f for f in data["items"] if f["field_type"] == "MATRIX")
|
||||
saved_scope = data["scope"]
|
||||
return "prepare_research_backtest", {
|
||||
"name": "Chatbox 数据集研究",
|
||||
"hypothesis": "验证所选合成字段的横截面排序信号",
|
||||
"template_input_id": data["id"],
|
||||
"input_snapshot_id": data["id"],
|
||||
"candidates": [
|
||||
{
|
||||
"client_item_id": "research-1",
|
||||
|
||||
@@ -30,7 +30,7 @@ async def acceptance():
|
||||
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.models import TemplateInput
|
||||
from app.models import ResearchInputSnapshot
|
||||
from app.research.workspace_contracts import TemplateSpec
|
||||
from tests.test_ai import configure
|
||||
from tests.test_backtests import setup
|
||||
@@ -66,7 +66,7 @@ async def acceptance():
|
||||
}
|
||||
|
||||
async with app.state.sessions() as db:
|
||||
fixed = await db.scalar(select(TemplateInput))
|
||||
fixed = await db.scalar(select(ResearchInputSnapshot))
|
||||
body = {
|
||||
"request_id": "finite-run",
|
||||
"name": "PG 有限研究",
|
||||
|
||||
@@ -170,14 +170,15 @@ async def main(args):
|
||||
fields = await api("GET", f"/catalog/datasets/pv1/fields?{params}&limit=100")
|
||||
fixed_input = await api(
|
||||
"POST",
|
||||
"/catalog/inputs",
|
||||
"/data-preparations/from-dataset",
|
||||
{
|
||||
"scope": scope,
|
||||
"dataset_id": "pv1",
|
||||
"collection_version": fields["collection_version"],
|
||||
"selection": "all",
|
||||
},
|
||||
)
|
||||
fixed_input = (await api("POST", "/data-preparations/freeze", {
|
||||
"items": [{"id": fixed_input["id"], "version": fixed_input["version"]}]}))["items"][0]
|
||||
availability = await api(
|
||||
"POST", "/catalog/field-availability/refresh", {"field_id": "close", "scope": scope}
|
||||
)
|
||||
|
||||
@@ -28,7 +28,7 @@ async def acceptance():
|
||||
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.models import ResearchExperiment, ResearchParent, TemplateInput
|
||||
from app.models import ResearchExperiment, ResearchInputSnapshot, ResearchParent
|
||||
from tests.test_research_outcomes import (
|
||||
test_feature_conversion_keeps_original_version_through_experiment,
|
||||
test_lineage_retains_multiple_parents_and_descendants,
|
||||
@@ -46,7 +46,7 @@ async def acceptance():
|
||||
)
|
||||
assert response.status_code == 200
|
||||
async with app.state.sessions() as db:
|
||||
fixed = await db.scalar(select(TemplateInput))
|
||||
fixed = await db.scalar(select(ResearchInputSnapshot))
|
||||
for experiment in await db.scalars(select(ResearchExperiment)):
|
||||
for parent in experiment.parents:
|
||||
assert await db.get(ResearchParent, (experiment.id, parent["kind"], parent["id"]))
|
||||
|
||||
@@ -30,7 +30,7 @@ async def acceptance():
|
||||
|
||||
from app.config import Settings
|
||||
from app.main import create_app
|
||||
from app.models import TemplateInput
|
||||
from app.models import ResearchInputSnapshot
|
||||
from app.research.workspace_contracts import TemplateSpec
|
||||
from tests.test_ai import configure
|
||||
from tests.test_backtests import setup
|
||||
@@ -67,7 +67,7 @@ async def acceptance():
|
||||
}
|
||||
|
||||
async with app.state.sessions() as db:
|
||||
fixed = await db.scalar(select(TemplateInput))
|
||||
fixed = await db.scalar(select(ResearchInputSnapshot))
|
||||
body = {
|
||||
"request_id": "finite-run",
|
||||
"name": "PG 有限研究",
|
||||
|
||||
@@ -10,11 +10,10 @@ from sqlalchemy import func, select
|
||||
from app.ai.capabilities import ToolContext, assemble
|
||||
from app.ai.tools import CAPABILITIES
|
||||
from app.business import Business
|
||||
from app.models import AIMessage, AIRun, AIToolCall, Alpha, Research, TemplateInput
|
||||
from app.models import AIMessage, AIRun, AIToolCall, Alpha, Research, ResearchInputSnapshot
|
||||
from app.research.service import ResearchBuilder
|
||||
from tests.test_ai import configure, single_tool_factory, start
|
||||
from tests.test_api import seed
|
||||
from tests.test_catalog import SCOPE
|
||||
from tests.test_catalog import catalog as catalog_fixture
|
||||
from tests.test_research_integration import fixed_input as fixed_input_fixture
|
||||
|
||||
@@ -62,15 +61,13 @@ async def test_prepare_failure_rolls_back_artifact_but_keeps_audit(app, logged_i
|
||||
# select_input has already persisted the new input before requesting its result page.
|
||||
raise HTTPException(422, "准备输入后的校验失败")
|
||||
|
||||
updated = await logged_in.patch(f"/api/v1/data-preparations/{fixed_input['preparation_id']}",
|
||||
json={"version": fixed_input["preparation_version"], "name": "new version"})
|
||||
assert updated.status_code == 200
|
||||
monkeypatch.setattr(ResearchBuilder, "input_page", unavailable_page)
|
||||
app.state.ai.model_factory = single_tool_factory(
|
||||
"prepare_research_input",
|
||||
{
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
"collection_version": fixed_input["collection_version"],
|
||||
"field_ids": ["TEST_FIN_001"],
|
||||
},
|
||||
{"items": [{"id": fixed_input["preparation_id"], "version": updated.json()["version"]}]},
|
||||
)
|
||||
_, run, _ = await start(app, logged_in, "保存研究输入")
|
||||
call = run["tools"][0]
|
||||
@@ -78,7 +75,7 @@ async def test_prepare_failure_rolls_back_artifact_but_keeps_audit(app, logged_i
|
||||
assert call["presentation"]["effect"] == "prepare"
|
||||
assert call["result"]["error"] == "准备输入后的校验失败"
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(TemplateInput)) == 1
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 1
|
||||
assert (await db.get(AIToolCall, call["id"])).status == "failed"
|
||||
|
||||
|
||||
|
||||
@@ -212,8 +212,7 @@ def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, m
|
||||
},
|
||||
)
|
||||
for _ in range(2):
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
command.upgrade(config, "0014")
|
||||
alphas = sa.Table("alphas", sa.MetaData(), autoload_with=engine)
|
||||
with engine.connect() as db:
|
||||
rows = db.execute(sa.select(alphas).order_by(alphas.c.id)).mappings().all()
|
||||
@@ -224,4 +223,6 @@ def test_migration_backfills_multiple_batches_and_preserves_research(tmp_path, m
|
||||
record = db.execute(sa.select(research)).mappings().one()
|
||||
assert record["tags"] == ["PPAC"] and record["note"] == "keep" and record["version"] == 7
|
||||
command.downgrade(config, "0010")
|
||||
command.upgrade(config, "head")
|
||||
command.check(config)
|
||||
engine.dispose()
|
||||
|
||||
@@ -203,6 +203,7 @@ async def test_export_formula_injection_and_detail_variants(app, logged_in):
|
||||
assert row["name"] == "'=DANGEROUS()" and row["note"] == "' @formula()"
|
||||
assert (await logged_in.get(f"{PREFIX}/alphas/super1/pnl")).json() == {
|
||||
"cached": False,
|
||||
"series": [],
|
||||
"points": [],
|
||||
"fetched_at": None,
|
||||
}
|
||||
|
||||
@@ -86,16 +86,24 @@ async def search(client, suffix="/datasets", **params):
|
||||
|
||||
|
||||
async def prepare(client, version, **changes):
|
||||
return await client.post(
|
||||
BASE + "/inputs",
|
||||
json={
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
response = await client.post("/api/v1/data-preparations/from-dataset", json={
|
||||
"scope": changes.get("scope", SCOPE), "dataset_id": changes.get("dataset_id", "TEST_FIN"),
|
||||
"collection_version": version,
|
||||
"selection": "all",
|
||||
**changes,
|
||||
},
|
||||
)
|
||||
})
|
||||
if response.status_code != 201:
|
||||
return response
|
||||
collection = response.json()
|
||||
if changes.get("excluded_ids"):
|
||||
response = await client.patch(f"/api/v1/data-preparations/{collection['id']}/fields", json={
|
||||
"version": collection["version"], "remove_ids": changes["excluded_ids"]})
|
||||
if response.status_code != 200:
|
||||
return response
|
||||
collection = response.json()
|
||||
response = await client.post("/api/v1/data-preparations/freeze", json={
|
||||
"items": [{"id": collection["id"], "version": collection["version"]}]})
|
||||
if response.status_code != 201:
|
||||
return response
|
||||
return httpx.Response(201, json=response.json()["items"][0])
|
||||
|
||||
|
||||
async def test_complete_workflow_filters_notes_immutable_input(catalog):
|
||||
@@ -112,7 +120,7 @@ async def test_complete_workflow_filters_notes_immutable_input(catalog):
|
||||
response = await prepare(client, version)
|
||||
assert response.status_code == 201, response.text
|
||||
draft = response.json()
|
||||
assert len(draft["field_ids"]) == 123 and draft["status"] == "draft"
|
||||
assert len(draft["field_ids"]) == 123 and draft["preparation_version"] == 1
|
||||
for suffix in ["/datasets/TEST_FIN", "/datasets/TEST_FIN/fields/TEST_FIN_122"]:
|
||||
detail = await search(client, suffix)
|
||||
assert detail["research"]["version"] == 1
|
||||
@@ -135,7 +143,7 @@ async def test_complete_workflow_filters_notes_immutable_input(catalog):
|
||||
assert newer["collection_version"] != version
|
||||
assert (await prepare(client, version)).status_code == 409
|
||||
assert len((await prepare(client, newer["collection_version"])).json()["field_ids"]) == 125
|
||||
assert (await client.get(BASE + "/inputs/" + draft["id"])).json() == draft
|
||||
assert (await client.get("/api/v1/research/input-snapshots/" + draft["id"])).json() == draft
|
||||
assert (await search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122"))["research"][
|
||||
"note"
|
||||
] == "保留研究备注"
|
||||
@@ -219,7 +227,7 @@ async def test_scope_ownership_empty_and_unknown_fields_are_rejected(catalog):
|
||||
await sync(catalog, scope=other)
|
||||
await sync(catalog, "TEST_FIN", scope=other)
|
||||
assert (await prepare(client, version, scope=other)).status_code == 409
|
||||
assert len((await client.get(BASE + "/inputs", params=SCOPE)).json()) == 1
|
||||
assert (await client.get(BASE + "/inputs", params=SCOPE)).status_code == 404
|
||||
|
||||
|
||||
async def test_catalog_authentication_and_origin(app, client):
|
||||
|
||||
@@ -162,7 +162,8 @@ async def test_official_sdk_client_and_error_contract(mcp_app):
|
||||
async with ClientSession(streams[0], streams[1]) as client:
|
||||
await client.initialize()
|
||||
listed = await client.list_tools()
|
||||
assert len(listed.tools) == 18
|
||||
assert len(listed.tools) == 20
|
||||
assert {"search_data_preparations", "get_data_preparation"} <= {t.name for t in listed.tools}
|
||||
caps = await client.call_tool("get_research_capabilities", {})
|
||||
assert caps.structured_content["max_candidates"] == 100
|
||||
result = await client.call_tool("submit_backtests", submission())
|
||||
@@ -459,3 +460,27 @@ async def test_worldquant_authentication_permissions_and_challenge(mcp_app):
|
||||
assert missing.structured_content["error"]["code"] == "CREDENTIALS_NOT_CONFIGURED"
|
||||
wrong = await mcp_app.state.mcp.invoke(principal, "get_worldquant_connection", {"job_id": "unrelated"})
|
||||
assert wrong.structured_content["error"]["code"] == "NOT_FOUND"
|
||||
|
||||
async def test_mcp_preparations_freeze_at_submit_and_survive_deletion(mcp_app):
|
||||
from app.catalog.contracts import Scope
|
||||
from app.preparations.contracts import PreparationReference
|
||||
from app.preparations.service import Preparations
|
||||
principal, _ = await credentials(mcp_app)
|
||||
async with mcp_app.state.sessions.begin() as db:
|
||||
collection = await Preparations(db).create("MCP prepared", "", Scope(region="USA", universe="TOP3000", delay=1),
|
||||
[{"id": "close", "field_id": "close", "name": "Close", "dataset_id": "pv1", "dataset_name": "Price",
|
||||
"description": "Synthetic close", "field_type": "MATRIX", "source": "local", "fetched_at": now().isoformat()}])
|
||||
refs = [{"id": collection["id"], "version": 1}]
|
||||
found = await invoke(mcp_app, principal, "search_data_preparations", {"q": "MCP prepared"})
|
||||
assert found["items"][0]["id"] == collection["id"]
|
||||
detail = await invoke(mcp_app, principal, "get_data_preparation", {**refs[0], "limit": 1})
|
||||
assert detail["fields"]["items"][0]["dataset_id"] == "pv1"
|
||||
result = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
|
||||
async with mcp_app.state.sessions.begin() as db:
|
||||
run = await db.get(BacktestRun, result["backtest_run_id"])
|
||||
snapshot_id = run.source["input_snapshot_ids"][0]
|
||||
await Preparations(db).remove([PreparationReference(**refs[0])])
|
||||
assert (await Preparations(db).snapshot(snapshot_id))["fields"][0]["description"] == "Synthetic close"
|
||||
# Idempotent replay uses the already fixed run even after the collection is gone.
|
||||
replay = await invoke(mcp_app, principal, "submit_backtests", submission(preparation_refs=refs))
|
||||
assert replay["backtest_run_id"] == result["backtest_run_id"]
|
||||
|
||||
@@ -0,0 +1,352 @@
|
||||
"""Preparation contracts across real HTTP, transactions and immutable research inputs."""
|
||||
|
||||
from argparse import Namespace
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.catalog.contracts import CatalogJobInput, Scope
|
||||
from app.catalog.service import Catalog
|
||||
from app.models import CatalogBatch, Job, ResearchInputSnapshot
|
||||
from tests.catalog_fake import field_records
|
||||
from tests.test_catalog import SCOPE, sync
|
||||
from tests.test_catalog import catalog as catalog_fixture
|
||||
|
||||
catalog = catalog_fixture
|
||||
BASE = "/api/v1/data-preparations"
|
||||
|
||||
|
||||
def reference(field):
|
||||
return {k: field[k] for k in ("scope", "dataset_id", "field_id", "source", "collection_version")}
|
||||
|
||||
|
||||
async def copied(catalog):
|
||||
client = catalog[0]
|
||||
await sync(catalog)
|
||||
job = await sync(catalog, "TEST_FIN")
|
||||
response = await client.post(
|
||||
BASE + "/from-dataset",
|
||||
json={"scope": SCOPE, "dataset_id": "TEST_FIN", "collection_version": job["id"]},
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
return response.json()
|
||||
|
||||
|
||||
async def test_multi_dataset_collection_atomic_changes_and_snapshot_independence(catalog):
|
||||
client, runner, state = catalog
|
||||
collection = await copied(catalog)
|
||||
state["fields"] = field_records("TEST_NEWS", 3)
|
||||
await sync(catalog, "TEST_NEWS")
|
||||
query = (await client.get("/api/v1/catalog/fields", params={**SCOPE, "dataset_id": "TEST_NEWS"})).json()
|
||||
assert query["total"] == 3
|
||||
ref = reference(query["items"][0])
|
||||
response = await client.patch(
|
||||
f"{BASE}/{collection['id']}/fields", json={"version": 1, "fields": [ref, ref]}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
collection = response.json()
|
||||
assert collection["field_count"] == 124 and collection["dataset_count"] == 2
|
||||
bad = {**ref, "scope": {**SCOPE, "delay": 0}}
|
||||
response = await client.patch(
|
||||
f"{BASE}/{collection['id']}/fields",
|
||||
json={"version": 2, "fields": [ref, bad], "remove_ids": ["TEST_FIN_001"]},
|
||||
)
|
||||
assert response.status_code == 422
|
||||
assert (await client.get(f"{BASE}/{collection['id']}")).json()["version"] == 2
|
||||
snapshot = (
|
||||
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
|
||||
).json()["items"][0]
|
||||
assert snapshot["dataset_ids"] == ["TEST_FIN", "TEST_NEWS"]
|
||||
assert all("description" in f and f["dataset_id"] for f in snapshot["fields"])
|
||||
assert (
|
||||
await client.patch(f"{BASE}/{collection['id']}", json={"version": 2, "name": "edited"})
|
||||
).status_code == 200
|
||||
assert (
|
||||
await client.post(BASE + "/freeze", json={"items": [{"id": collection["id"], "version": 2}]})
|
||||
).status_code == 409
|
||||
assert (await client.delete(f"{BASE}/{collection['id']}?version=3")).status_code == 200
|
||||
assert (await client.get(f"{BASE}/{collection['id']}")).status_code == 404
|
||||
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot['id']}")).json() == snapshot
|
||||
async with runner.sessions() as db:
|
||||
assert await db.get(ResearchInputSnapshot, snapshot["id"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key,value", [("region", "CHN"), ("universe", "TOP1000"), ("delay", 0), ("instrument_type", "FUTURE")]
|
||||
)
|
||||
async def test_scope_dimensions_rejected_without_creating_collection(catalog, key, value):
|
||||
client = catalog[0]
|
||||
await copied(catalog)
|
||||
field = (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["items"][0]
|
||||
before = (await client.get(BASE)).json()["total"]
|
||||
response = await client.post(
|
||||
BASE, json={"name": "bad", "scope": {**SCOPE, key: value}, "fields": [reference(field)]}
|
||||
)
|
||||
assert response.status_code == 422
|
||||
assert (await client.get(BASE)).json()["total"] == before
|
||||
|
||||
|
||||
async def test_online_fields_without_sync_and_local_search_before_pagination(catalog):
|
||||
client, _, _ = catalog
|
||||
response = await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100})
|
||||
assert response.status_code == 200, response.text
|
||||
first = response.json()
|
||||
assert len(first["items"]) == 100 and first["has_more"]
|
||||
second = (
|
||||
await client.get("/api/v1/catalog/worldquant/fields", params={**SCOPE, "limit": 100, "offset": 100})
|
||||
).json()
|
||||
assert len(second["items"]) > 0
|
||||
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
|
||||
field = reference(second["items"][-1])
|
||||
response = await client.post(BASE, json={"name": "online only", "scope": SCOPE, "fields": [field]})
|
||||
assert response.status_code == 201, response.text
|
||||
assert response.json()["field_count"] == 1
|
||||
assert (await client.get("/api/v1/catalog/fields", params=SCOPE)).json()["total"] == 0
|
||||
await copied(catalog)
|
||||
query = (
|
||||
await client.get(
|
||||
"/api/v1/catalog/fields", params={**SCOPE, "q": "合成字段说明 11", "limit": 2, "offset": 2}
|
||||
)
|
||||
).json()
|
||||
assert query["total"] == 11 and len(query["items"]) == 2
|
||||
response = await client.get(
|
||||
"/api/v1/catalog/fields", params={**SCOPE, "coverage_min": 0.9, "coverage_max": 0.5}
|
||||
)
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
async def test_empty_crud_copy_and_batch_delete_are_version_checked(catalog):
|
||||
client = catalog[0]
|
||||
first = (await client.post(BASE, json={"name": "empty", "scope": SCOPE})).json()
|
||||
assert first["field_count"] == 0
|
||||
assert (
|
||||
await client.post(BASE + "/freeze", json={"items": [{"id": first["id"], "version": 1}]})
|
||||
).status_code == 422
|
||||
second = (await client.post(f"{BASE}/{first['id']}/copy", json={"version": 1})).json()
|
||||
response = await client.post(
|
||||
BASE + "/batch-delete",
|
||||
json={"items": [{"id": first["id"], "version": 1}, {"id": second["id"], "version": 99}]},
|
||||
)
|
||||
assert response.status_code == 409
|
||||
assert (await client.get(BASE)).json()["total"] == 2
|
||||
assert (
|
||||
await client.patch(
|
||||
f"{BASE}/{first['id']}", json={"name": "x", "version": 1, "scope": {**SCOPE, "delay": 0}}
|
||||
)
|
||||
).status_code == 422
|
||||
|
||||
|
||||
async def test_full_sync_partial_failure_and_resume_publish_per_dataset(catalog):
|
||||
client, runner, state = catalog
|
||||
await copied(catalog)
|
||||
async with runner.sessions.begin() as db:
|
||||
service = Catalog(db)
|
||||
job = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
duplicate = await service.create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
assert duplicate.id == job.id
|
||||
# The fixture deliberately returns TEST_FIN ownership for the other datasets.
|
||||
await runner.execute(job.id)
|
||||
failed = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert failed["status"] == "completed_with_errors" and failed["processed"] == 1
|
||||
version = (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
|
||||
"collection_version"
|
||||
]
|
||||
assert version != job.id
|
||||
async with runner.sessions() as db:
|
||||
assert (await db.get(CatalogBatch, version)).job_id == job.id
|
||||
state["fields"] = field_records("TEST_NEWS", 3)
|
||||
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
|
||||
await runner.execute(job.id)
|
||||
retried = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert retried["processed"] == 2 and retried["failed"] == 1
|
||||
assert (await client.get("/api/v1/catalog/datasets/TEST_FIN", params=SCOPE)).json()[
|
||||
"collection_version"
|
||||
] == version
|
||||
state["fields"] = field_records("TEST_UNKNOWN", 0)
|
||||
await client.post(f"/api/v1/sync-jobs/{job.id}/retry")
|
||||
await runner.execute(job.id)
|
||||
done = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert done["status"] == "completed" and done["processed"] == 3 and done["failed"] == 0
|
||||
|
||||
|
||||
async def test_cli_timeout_reuses_active_job_and_does_not_cancel(catalog, monkeypatch, capsys):
|
||||
from app import cli
|
||||
|
||||
_, runner, _ = catalog
|
||||
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
|
||||
args = Namespace(
|
||||
region="USA",
|
||||
universe="TOP3000",
|
||||
delay=1,
|
||||
instrument_type="EQUITY",
|
||||
resume_job=None,
|
||||
wait_timeout=0.01,
|
||||
)
|
||||
assert await cli.catalog_sync_command(args) == 4
|
||||
assert await cli.catalog_sync_command(args) == 4
|
||||
async with runner.sessions() as db:
|
||||
rows = list(await db.scalars(select(Job).where(Job.kind == "catalog_full_sync")))
|
||||
assert len(rows) == 1 and rows[0].status == "queued" and not rows[0].cancel_requested
|
||||
assert "等待超时" in capsys.readouterr().out
|
||||
|
||||
|
||||
async def test_collection_version_checked_at_research_submission(catalog):
|
||||
client = catalog[0]
|
||||
collection = await copied(catalog)
|
||||
response = await client.patch(f"{BASE}/{collection['id']}", json={"version": 1, "name": "changed"})
|
||||
assert response.status_code == 200
|
||||
body = {
|
||||
"inline": {
|
||||
"name": "prepared backtest",
|
||||
"preparation_refs": [{"id": collection["id"], "version": 1}],
|
||||
"candidates": [
|
||||
{
|
||||
"client_item_id": "one",
|
||||
"expression": "rank(TEST_FIN_001)",
|
||||
"settings": {k: SCOPE[k] for k in ("region", "universe", "delay")},
|
||||
}
|
||||
],
|
||||
}
|
||||
}
|
||||
assert (await client.post("/api/v1/backtests/previews", json=body)).status_code == 409
|
||||
body["inline"]["preparation_refs"][0]["version"] = 2
|
||||
response = await client.post("/api/v1/backtests/previews", json=body)
|
||||
assert response.status_code == 201, response.text
|
||||
snapshot_id = response.json()["source"]["input_snapshot_ids"][0]
|
||||
await client.delete(f"{BASE}/{collection['id']}?version=2")
|
||||
assert (await client.get(f"/api/v1/research/input-snapshots/{snapshot_id}")).json()["fields"][1][
|
||||
"dataset_id"
|
||||
] == "TEST_FIN"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"status,code",
|
||||
[
|
||||
("completed", 0),
|
||||
("completed_with_errors", 1),
|
||||
("failed", 1),
|
||||
("cancelled", 1),
|
||||
("waiting_connection", 3),
|
||||
("waiting_auth", 3),
|
||||
],
|
||||
)
|
||||
async def test_cli_terminal_exit_codes(catalog, monkeypatch, status, code):
|
||||
from app import cli
|
||||
|
||||
_, runner, _ = catalog
|
||||
monkeypatch.setattr(cli, "Settings", lambda: runner.settings)
|
||||
original = Catalog.create_job
|
||||
|
||||
async def terminal(self, body, **kwargs):
|
||||
result = await original(self, body, **kwargs)
|
||||
row = await self.db.get(Job, result.id)
|
||||
row.status = status
|
||||
return result
|
||||
|
||||
monkeypatch.setattr(Catalog, "create_job", terminal)
|
||||
args = Namespace(
|
||||
region="USA", universe="TOP3000", delay=1, instrument_type="EQUITY", resume_job=None, wait_timeout=1
|
||||
)
|
||||
assert await cli.catalog_sync_command(args) == code
|
||||
|
||||
|
||||
async def test_full_sync_restart_keeps_page_and_auth_pauses_all_datasets(catalog, monkeypatch):
|
||||
import asyncio
|
||||
|
||||
from app.worldquant import WqError
|
||||
|
||||
client, runner, _ = catalog
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
original = runner.client.catalog_page
|
||||
interrupted = False
|
||||
|
||||
async def pages(scope, dataset_id, offset):
|
||||
nonlocal interrupted
|
||||
if dataset_id == "TEST_FIN" and offset == 100 and not interrupted:
|
||||
interrupted = True
|
||||
runner.stopping = True
|
||||
raise asyncio.CancelledError()
|
||||
if dataset_id == "TEST_NEWS":
|
||||
raise WqError("synthetic network interruption", "network_error")
|
||||
return await original(scope, dataset_id, offset)
|
||||
|
||||
monkeypatch.setattr(runner.client, "catalog_page", pages)
|
||||
await runner.execute(job.id)
|
||||
async with runner.sessions() as db:
|
||||
batch = await db.scalar(
|
||||
select(CatalogBatch).where(CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_FIN")
|
||||
)
|
||||
assert batch.offset == 100 and not batch.complete
|
||||
assert (await db.get(Job, job.id)).status == "queued"
|
||||
runner.stopping = False
|
||||
await runner.execute(job.id)
|
||||
row = (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()
|
||||
assert row["status"] == "waiting_connection" and row["processed"] == 1
|
||||
assert row["checkpoint"]["dataset_id"] == "TEST_NEWS"
|
||||
assert row["checkpoint"]["datasets_completed"] == 1
|
||||
async with runner.sessions() as db:
|
||||
assert not await db.scalar(
|
||||
select(CatalogBatch).where(
|
||||
CatalogBatch.job_id == job.id, CatalogBatch.dataset_id == "TEST_UNKNOWN"
|
||||
)
|
||||
)
|
||||
await client.post(f"/api/v1/sync-jobs/{job.id}/cancel")
|
||||
assert (await client.get(f"/api/v1/sync-jobs/{job.id}")).json()["status"] == "cancelled"
|
||||
|
||||
|
||||
async def test_online_instrument_type_is_verified(catalog):
|
||||
client, _, state = catalog
|
||||
state["fields"][0]["instrumentType"] = "FUTURE"
|
||||
assert (await client.get("/api/v1/catalog/worldquant/fields", params=SCOPE)).status_code == 502
|
||||
|
||||
|
||||
async def test_mcp_collection_reads_and_versioned_submit_contract(catalog):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from app.research_access.contracts import PreparationRead, PreparationSearch, Submit
|
||||
from app.research_access.service import ResearchAccess
|
||||
|
||||
client, runner, _ = catalog
|
||||
collection = await copied(catalog)
|
||||
async with runner.sessions.begin() as db:
|
||||
access = ResearchAccess(
|
||||
db, SimpleNamespace(token_id="test", admin_id=1), runner.client, "http://testserver"
|
||||
)
|
||||
result = await access.preparations(PreparationSearch(q=collection["name"]))
|
||||
assert result["items"][0]["version"] == 1
|
||||
detail = await access.preparation(PreparationRead(id=collection["id"], version=1, limit=1))
|
||||
assert detail["fields"]["has_more"] and detail["fields"]["items"][0]["dataset_id"] == "TEST_FIN"
|
||||
# Schema accepts versioned references and rejects accidental old dataset input payloads.
|
||||
from tests.test_mcp import submission
|
||||
|
||||
payload = submission()
|
||||
payload["preparation_refs"] = [{"id": collection["id"], "version": 1}]
|
||||
assert Submit.model_validate(payload).preparation_refs[0].id == collection["id"]
|
||||
|
||||
async def test_retry_full_job_reuses_another_active_job(catalog):
|
||||
from app.business import Business
|
||||
_, runner, _ = catalog
|
||||
body = CatalogJobInput(scope=Scope(**SCOPE))
|
||||
async with runner.sessions.begin() as db:
|
||||
old = await Catalog(db).create_job(body, full=True)
|
||||
(await db.get(Job, old.id)).status = "failed"
|
||||
async with runner.sessions.begin() as db:
|
||||
new = await Catalog(db).create_job(body, full=True)
|
||||
assert new.id != old.id
|
||||
async with runner.sessions.begin() as db:
|
||||
result = await Business(db).retry_job(old.id)
|
||||
assert result["id"] == new.id
|
||||
assert (await db.get(Job, old.id)).status == "failed"
|
||||
|
||||
|
||||
async def test_retry_waiting_full_job_requeues_its_checkpoint(catalog):
|
||||
from app.business import Business
|
||||
_, runner, _ = catalog
|
||||
async with runner.sessions.begin() as db:
|
||||
job = await Catalog(db).create_job(CatalogJobInput(scope=Scope(**SCOPE)), full=True)
|
||||
row = await db.get(Job, job.id)
|
||||
row.status, row.checkpoint = "waiting_connection", {"offset": 100}
|
||||
async with runner.sessions.begin() as db:
|
||||
result = await Business(db).retry_job(job.id)
|
||||
assert result["status"] == "queued" and result["checkpoint"]["offset"] == 100
|
||||
@@ -11,7 +11,7 @@ from app.ai.tools import CAPABILITIES
|
||||
from app.alphas import upsert_alpha
|
||||
from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput
|
||||
from app.business import Business
|
||||
from app.models import BacktestPreview, BacktestRun, Research, TemplateInput
|
||||
from app.models import BacktestPreview, BacktestRun, Research, ResearchInputSnapshot
|
||||
from tests.test_ai import configure, single_tool_factory
|
||||
from tests.test_backtests import execute, setup, start
|
||||
from tests.test_catalog import SCOPE, prepare, sync
|
||||
@@ -47,7 +47,7 @@ def construction(input_id):
|
||||
return {
|
||||
"name": "字段研究",
|
||||
"hypothesis": "显式字段排序",
|
||||
"template_input_id": input_id,
|
||||
"input_snapshot_id": input_id,
|
||||
"candidates": [
|
||||
{
|
||||
"client_item_id": "one",
|
||||
@@ -66,7 +66,7 @@ async def test_chatbox_catalog_to_results_and_alpha_sources(app, logged_in, fixe
|
||||
conversation = (await logged_in.post("/api/v1/ai/conversations")).json()["id"]
|
||||
context = {"page": "datasets", "catalog_scope": SCOPE, "dataset_id": "TEST_FIN"}
|
||||
if use_saved_input:
|
||||
context["template_input_id"] = fixed_input["id"]
|
||||
context["input_snapshot_id"] = fixed_input["id"]
|
||||
run = await ask(logged_in, conversation, "研究此输入" if use_saved_input else "自行选字段研究", context)
|
||||
assert run["status"] == "waiting_approval", run
|
||||
assert not platform.posts
|
||||
@@ -75,7 +75,7 @@ async def test_chatbox_catalog_to_results_and_alpha_sources(app, logged_in, fixe
|
||||
assert source["kind"] == "chatbox"
|
||||
assert source["reference"] == conversation
|
||||
assert source["research_id"] == run["id"]
|
||||
assert source["template_input_id"]
|
||||
assert source["input_snapshot_id"]
|
||||
assert approval["preview"]["backtest"]["items"][0]["expression"] == "rank(TEST_FIN_001)"
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(BacktestRun)) == 0
|
||||
@@ -128,7 +128,7 @@ async def test_invalid_construction_has_no_partial_preview(logged_in, app, fixed
|
||||
elif invalid == "unknown_type":
|
||||
item["bindings"]["signal"] = {"field_id": "TEST_FIN_122", "field_type": "MATRIX"}
|
||||
else:
|
||||
body["template_input_id"] = "missing"
|
||||
body["input_snapshot_id"] = "missing"
|
||||
response = await logged_in.post("/api/v1/backtests/research-previews", json=body)
|
||||
assert response.status_code == (404 if invalid == "missing_input" else 422), response.text
|
||||
async with app.state.sessions() as db:
|
||||
@@ -144,15 +144,10 @@ async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in
|
||||
assert page["field_count"] == page["total"] == 123 and len(page["items"]) == 23
|
||||
assert page["items"][-1]["field_type"] == "FUTURE_TYPE" and not page["has_more"]
|
||||
assert page["_meta"]["source"] == "local_database"
|
||||
selected = await tool(
|
||||
"prepare_research_input",
|
||||
{
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
"collection_version": fixed_input["collection_version"],
|
||||
"field_ids": ["TEST_FIN_001"],
|
||||
},
|
||||
)
|
||||
subset = (await logged_in.post("/api/v1/data-preparations", json={"name": "one field", "scope": SCOPE,
|
||||
"fields": [{"scope": SCOPE, "dataset_id": "TEST_FIN", "field_id": "TEST_FIN_001", "source": "local",
|
||||
"collection_version": fixed_input["fields"][0]["collection_version"]}]})).json()
|
||||
selected = await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
|
||||
assert selected["field_count"] == 1
|
||||
bad = construction(selected["id"])
|
||||
bad["candidates"][0]["bindings"]["signal"]["field_id"] = "TEST_FIN_002"
|
||||
@@ -166,19 +161,12 @@ async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in
|
||||
"/api/v1/backtests/research-previews", json=construction(fixed_input["id"])
|
||||
)
|
||||
assert response.status_code == 201, response.text
|
||||
await logged_in.patch(f"/api/v1/data-preparations/{subset['id']}", json={"name": "edited", "version": subset["version"]})
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await tool(
|
||||
"prepare_research_input",
|
||||
{
|
||||
"scope": SCOPE,
|
||||
"dataset_id": "TEST_FIN",
|
||||
"collection_version": fixed_input["collection_version"],
|
||||
"field_ids": ["TEST_FIN_001"],
|
||||
},
|
||||
)
|
||||
await tool("prepare_research_input", {"items": [{"id": subset["id"], "version": subset["version"]}]})
|
||||
assert exc.value.status_code == 409
|
||||
async with app.state.sessions() as db:
|
||||
assert await db.scalar(select(func.count()).select_from(TemplateInput)) == 2
|
||||
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 2
|
||||
|
||||
|
||||
async def test_multiple_origins_preserve_research_and_do_not_duplicate_alphas(app, logged_in):
|
||||
|
||||
@@ -213,7 +213,7 @@ async def test_operator_annotation_and_refresh_preserve_local(app, logged_in, re
|
||||
|
||||
|
||||
async def test_settings_variant_requires_all_fields_in_target(app, logged_in, research_input, catalog):
|
||||
from app.models import CatalogScope, TemplateInput
|
||||
from app.models import CatalogScope, ResearchInputSnapshot
|
||||
from app.research.experiments import Experiments
|
||||
from app.research.workspace_contracts import SettingVariants
|
||||
|
||||
@@ -232,14 +232,13 @@ async def test_settings_variant_requires_all_fields_in_target(app, logged_in, re
|
||||
db.add(CatalogScope(key=target_key, scope=target_scope))
|
||||
await db.flush()
|
||||
db.add(
|
||||
TemplateInput(
|
||||
ResearchInputSnapshot(
|
||||
id="target",
|
||||
scope_key=target_key,
|
||||
dataset_id="TEST_FIN",
|
||||
collection_version=research_input["collection_version"],
|
||||
selection="explicit",
|
||||
field_ids=["TEST_FIN_001"],
|
||||
field_types={"TEST_FIN_001": "MATRIX"},
|
||||
preparation_id="target-preparation",
|
||||
preparation_version=1,
|
||||
content={**{k: v for k, v in research_input.items() if k not in ("id", "preparation_id", "preparation_version", "created_at")}, "scope": target_scope,
|
||||
"field_ids": ["TEST_FIN_001"], "field_types": {"TEST_FIN_001": "MATRIX"},
|
||||
"fields": [f for f in research_input["fields"] if f["id"] == "TEST_FIN_001"]},
|
||||
)
|
||||
)
|
||||
await db.flush()
|
||||
@@ -514,3 +513,16 @@ def test_real_seed_settings_preserve_execution_options_and_reject_unknowns():
|
||||
assert snapshot["startDate"] == "2014-01-01"
|
||||
with pytest.raises(ValidationError):
|
||||
seed_settings({**snapshot, "unknownOption": True})
|
||||
|
||||
async def test_model_cannot_append_preparation_references_to_fixed_inputs(app, logged_in, research_input, monkeypatch):
|
||||
from app.models import ResearchAsset
|
||||
from app.research import routes
|
||||
from app.research.workspace_contracts import FeatureSpec
|
||||
async def model(*args):
|
||||
return FeatureSpec(name="untrusted", hypothesis="test", input_ids=[research_input["id"]],
|
||||
preparation_refs=[{"id": research_input["preparation_id"], "version": research_input["preparation_version"]}]), {}
|
||||
monkeypatch.setattr(routes, "request_model", model)
|
||||
response = await logged_in.post("/api/v1/research/generate", json={"name": "test", "hypothesis": "test", "method": "feature", "input_ids": [research_input["id"]]})
|
||||
assert response.status_code == 422 and "不能改变" in response.text
|
||||
async with app.state.sessions() as db:
|
||||
assert not await db.scalar(select(ResearchAsset).where(ResearchAsset.name == "untrusted"))
|
||||
|
||||
@@ -87,3 +87,30 @@ docker logs --tail=100 wq-alpha-production-web-1
|
||||
采用 compound 经验 1 中的生产/开发隔离、显式外部网络、必填凭据、固定项目名、独立迁移和部署健康检查。未照搬旧项目端口、业务 Job 或卷权限初始化:本应用后端不写 named volume,已有镜像使用 UID 10001。
|
||||
|
||||
凭据存放在 Gitea Secrets,通过步骤级环境变量交给 Compose,不落地到配置文件。Docker 容器仍需持有运行时配置,因此拥有 Runner 或 Docker 管理权限的人仍可能读取它们。若此前手动创建了旧 `.env.production`,新流程不再读取它;确认配置已迁移到 Gitea 并妥善备份密钥后可自行移除旧文件。相关官方资料:[Compose 外部网络](https://docs.docker.com/reference/compose-file/networks/)、[环境变量插值](https://docs.docker.com/compose/how-tos/environment-variables/variable-interpolation/)、[Gitea Runner 标签](https://gitea.com/gitea/runner/src/branch/main/README.md)。
|
||||
|
||||
## 5. 1Panel 夜间全量目录同步
|
||||
|
||||
在 1Panel 的计划任务中建立 Shell 脚本任务,执行周期由 1Panel 设置,例如每天深夜执行。每次命令明确指定一个范围:
|
||||
|
||||
```bash
|
||||
docker exec wq-alpha-production-backend-1 \
|
||||
python -m app.cli catalog-sync \
|
||||
--region USA --universe TOP3000 --delay 1
|
||||
```
|
||||
|
||||
`--instrument-type` 默认 `EQUITY`,`--wait-timeout` 默认 `21600` 秒(六小时)。容器已持有部署时注入的环境,定时脚本无需再次填写数据库、密码或加密密钥;不要加 `-it`。容器命令继承环境的行为见 [Docker exec 官方文档](https://docs.docker.com/reference/cli/docker/container/exec/)。启动前在网页完成 WorldQuant 连接,人工验证仍在网页处理。
|
||||
|
||||
CLI 只创建持久化任务并等待,不启动另一个同步执行器;后端必须正在运行。相同范围的活动全量任务复用同一任务 ID。任务先完整更新数据集清单,再逐个同步所有字段,每个数据集独立完整发布。普通单集失败保留上一版并继续其他数据集;鉴权和网络连接问题暂停全量任务。任务面板可查看阶段、当前数据集、字段分页位置、成功/失败统计,并取消或重试。
|
||||
|
||||
1Panel 执行日志保存 CLI 输出的任务 ID、状态变化、检查点和最终汇总。也可使用 `docker logs --tail=100 wq-alpha-production-backend-1` 排查执行器;命令不输出平台认证信息。不要公开包含研究数据的日志。
|
||||
|
||||
失败或部分失败时按日志中的任务 ID 从检查点继续,已成功发布的数据集不会重抓:
|
||||
|
||||
```bash
|
||||
docker exec wq-alpha-production-backend-1 \
|
||||
python -m app.cli catalog-sync --resume-job TASK_ID --wait-timeout 21600
|
||||
```
|
||||
|
||||
退出码:`0` 全部成功,`1` 失败/部分失败/取消,`2` 参数错误,`3` 需要连接或人工验证,`4` 等待超时。超时只结束 CLI 等待,后台任务继续;重复执行同范围命令可继续等待活动任务。恢复前先处理连接或验证问题。服务重启后由原执行器恢复持久化检查点;定时调度仅由 1Panel 负责。
|
||||
|
||||
真实平台协议和 1Panel 实际触发效果不属于模拟验收。
|
||||
|
||||
@@ -51,6 +51,8 @@ python -m app.cli mcp-token-revoke TOKEN_ID
|
||||
| authenticate_worldquant | `{action?:"connect"或"verify"}`;默认 connect,使用已保存凭据异步认证,返回 job_id;要求 research:refresh |
|
||||
| get_research_capabilities | `{}`,含单次候选上限及完整候选 schema |
|
||||
| create_research_template | `{template,hypothesis,source_item_ids,idempotency_key,reference?}`;保存调用方生成的完整模板,要求 research:write |
|
||||
| search_data_preparations | `{q?,scope_key?,limit?,offset?}`;查询可编辑集合及当前版本 |
|
||||
| get_data_preparation | `{id,version,q?,limit?,offset?}`;按版本分页预览字段与数据集归属,版本冲突重新选择 |
|
||||
| search_catalog | `{filters:{region,universe,delay,...},dataset_id?}`;省略 dataset_id 查数据集,提供则查字段 |
|
||||
| get_research_metadata | `{query:{kind,...}}`;kind 为 scopes/settings/operators/field_availability |
|
||||
| refresh_research_data | `{query:{kind,...}}`;kind 为 catalog/operators/settings/field_availability/pnl |
|
||||
@@ -60,7 +62,7 @@ python -m app.cli mcp-token-revoke TOKEN_ID
|
||||
| check_self_correlation | `{alpha_ids:[...]}`,1–100 个已导入 Alpha ID;异步返回 job_id |
|
||||
| get_self_correlation | `{alpha_id}`;只读最新本地结果,含缓存和 stale 状态 |
|
||||
| search_backtests | 来源、reference、status、带时区起止时间、scope、q、候选精确匹配及分页 |
|
||||
| submit_backtests | `{name,candidates,idempotency_key,duplicate_policy?,source?}` |
|
||||
| submit_backtests | `{name,candidates,idempotency_key,preparation_refs?,duplicate_policy?,source?}` |
|
||||
| get_backtest | `{run_id,after?,event_limit?}`,after 为事件游标 |
|
||||
| get_backtest_results | `{run_id,item_ids?,limit?,offset?}` |
|
||||
| get_backtest_artifact | `{item_id,kind,limit?,offset?,date_from?,date_to?}`,kind 为 snapshot/pnl |
|
||||
@@ -188,3 +190,5 @@ MCP_TEST_DATABASE_URL=postgresql+asyncpg://USER:PASSWORD@127.0.0.1:PORT/wq_mcp_t
|
||||
```
|
||||
|
||||
该脚本执行迁移、并发提交/控制、重启重放、回退及重升级,不用于个人库或生产库。生产启用、真实平台兼容性、真实额度和客户端实际凭据配置仍需另行授权验证。本功能不会恢复任何定时研究。
|
||||
|
||||
使用数据准备集合时,`preparation_refs` 为最多 20 个 `{id,version}`。提交时核对版本、范围和字段,固定独立快照并保存到回测来源;空集合或版本冲突不会创建运行。后续编辑或删除集合不影响回测。无需旧输入草稿接口。
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
import { FieldDirectory } from "./preparations/FieldDirectory";
|
||||
import { DataPreparationPage } from "./preparations/DataPreparationPage";
|
||||
import { SnapshotDialog } from "./preparations/SnapshotDialog";
|
||||
import { useCallback, useEffect, useRef, useState } from "react";
|
||||
import {
|
||||
Badge,
|
||||
@@ -390,6 +393,15 @@ export default function App() {
|
||||
onTask={taskCreated}
|
||||
/>
|
||||
</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"}>
|
||||
<DatasetPage
|
||||
onContext={setDatasetContext}
|
||||
@@ -507,6 +519,7 @@ export default function App() {
|
||||
<ModelSettingsPanel />
|
||||
<WorkspacePreferences account={account} onChange={actionDone} />
|
||||
</SideSheet>
|
||||
<SnapshotDialog action={aiAction} />
|
||||
<JobPanel
|
||||
visible={showJobs && !(viewport < 1440 && chatOpen)}
|
||||
chatOffset={chatOffset}
|
||||
@@ -560,6 +573,8 @@ export default function App() {
|
||||
: { page: "variants" as const },
|
||||
alphas: alphaContext,
|
||||
datasets: datasetContext,
|
||||
fields: { page: "fields" as const },
|
||||
preparations: { page: "preparations" as const },
|
||||
backtests: backtestContext,
|
||||
account: { page: "account" as const },
|
||||
"mcp-keys": { page: "account" as const },
|
||||
|
||||
@@ -17,6 +17,8 @@ export type PageContext = {
|
||||
| "alphas"
|
||||
| "account"
|
||||
| "datasets"
|
||||
| "fields"
|
||||
| "preparations"
|
||||
| "backtests"
|
||||
| "operators"
|
||||
| "templates"
|
||||
@@ -36,7 +38,7 @@ export type PageContext = {
|
||||
dataset_id?: string;
|
||||
field_id?: string;
|
||||
collection_version?: string;
|
||||
template_input_id?: string;
|
||||
input_snapshot_id?: string;
|
||||
unsaved_field_selection?: boolean;
|
||||
backtest_run_id?: string;
|
||||
backtest_preview_id?: string;
|
||||
|
||||
@@ -12,7 +12,7 @@ export function contextReferences(context: PageContext): ResearchReference[] {
|
||||
add("scope", "范围", `${scope.region}/${scope.universe}/D${scope.delay}`);
|
||||
add("dataset", "数据集", context.dataset_id);
|
||||
add("field", "字段", context.field_id);
|
||||
add("input", "固定输入", context.template_input_id);
|
||||
add("input", "固定输入", context.input_snapshot_id);
|
||||
add("version", "集合版本", context.collection_version);
|
||||
add("asset", "研究素材", context.research_asset_id);
|
||||
add("experiment", "实验", context.research_experiment_id);
|
||||
@@ -29,7 +29,7 @@ export function researchPrompts(context: PageContext): string[] {
|
||||
case "datasets":
|
||||
return context.unsaved_field_selection
|
||||
? ["解释当前数据集的字段和适用场景"]
|
||||
: context.template_input_id
|
||||
: context.input_snapshot_id
|
||||
? [
|
||||
"分析这个固定输入中的字段,提出可验证的研究假设",
|
||||
"基于这个固定输入构建候选,并预览回测",
|
||||
@@ -63,6 +63,8 @@ const contextLabels: Record<
|
||||
PageContext["page"],
|
||||
(context: PageContext) => string
|
||||
> = {
|
||||
fields: () => "上下文:字段目录",
|
||||
preparations: () => "上下文:数据准备",
|
||||
home: () => "上下文:首页看板",
|
||||
alphas: (context) =>
|
||||
`上下文:${context.alpha_id ? `Alpha ${context.alpha_id}` : "Alpha 列表"}${context.selected_ids?.length ? ` · 已选 ${context.selected_ids.length} 条` : ""}`,
|
||||
@@ -78,7 +80,7 @@ const contextLabels: Record<
|
||||
account: () => "上下文:个人信息页",
|
||||
backtests: () => "上下文:回测研究",
|
||||
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 {
|
||||
@@ -103,7 +105,7 @@ const destinations: Record<UIAction["type"], Destination> = {
|
||||
open_experiment: { page: "templates", chat: "responsive" },
|
||||
open_variant: { page: "variants", chat: "responsive" },
|
||||
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_preview: { page: "backtests", chat: "responsive" },
|
||||
open_alpha: { page: "alphas", chat: "responsive" },
|
||||
|
||||
@@ -82,6 +82,7 @@ export const stateOptions = Object.entries(stateLabels).map(
|
||||
([value, label]) => ({ value, label }),
|
||||
);
|
||||
export const jobLabels: Record<string, string> = {
|
||||
catalog_full_sync: "全量同步数据目录与字段",
|
||||
catalog_sync: "同步数据集目录",
|
||||
field_sync: "同步数据字段",
|
||||
full_sync: "全量同步 Alpha",
|
||||
|
||||
@@ -1,3 +1,8 @@
|
||||
import {
|
||||
ResearchDataInput,
|
||||
researchSelection,
|
||||
} from "../preparations/ResearchDataInput";
|
||||
import type { InputSnapshot } from "../research/workspaceTypes";
|
||||
import { useCallback, useEffect, useState } from "react";
|
||||
import {
|
||||
Banner,
|
||||
@@ -70,6 +75,8 @@ export function BacktestPage({
|
||||
const [editor, setEditor] = useState(false);
|
||||
const [draft, setDraft] = useState<Draft | null>(null);
|
||||
const [name, setName] = useState("");
|
||||
const [inputIds, setInputIds] = useState<string[]>([]);
|
||||
const [inputs, setInputs] = useState<InputSnapshot[]>([]);
|
||||
const [source, setSource] = useState<Source>({ kind: "manual" });
|
||||
const [text, setText] = useState("");
|
||||
const [mode, setMode] = useState("lines");
|
||||
@@ -215,6 +222,8 @@ export function BacktestPage({
|
||||
setDraft(null);
|
||||
setName("");
|
||||
setSource({ kind: "manual" });
|
||||
setInputIds([]);
|
||||
setInputs([]);
|
||||
setText("");
|
||||
setSettings(initialSettings);
|
||||
setMode("lines");
|
||||
@@ -241,7 +250,7 @@ export function BacktestPage({
|
||||
}));
|
||||
if (!Array.isArray(candidates) || !candidates.length)
|
||||
throw new Error("请提供非空候选集合");
|
||||
return { name, source, candidates };
|
||||
return { name, source, candidates, ...researchSelection(inputIds, inputs) };
|
||||
}
|
||||
async function loadDraft(id: string) {
|
||||
setSettingsDraft(null);
|
||||
@@ -249,6 +258,8 @@ export function BacktestPage({
|
||||
setDraft(d);
|
||||
setName(d.name);
|
||||
setSource(d.source);
|
||||
setInputIds(d.source.input_snapshot_ids ?? []);
|
||||
setInputs([]);
|
||||
setMode("json");
|
||||
setText(JSON.stringify(d.candidates, null, 2));
|
||||
setPreview(null);
|
||||
@@ -600,6 +611,23 @@ export function BacktestPage({
|
||||
maxLength={200}
|
||||
/>
|
||||
</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">
|
||||
<label>
|
||||
来源
|
||||
|
||||
@@ -39,7 +39,8 @@ export type Source = {
|
||||
kind: string;
|
||||
reference?: string | null;
|
||||
batch_id?: string | null;
|
||||
template_input_id?: string | null;
|
||||
input_snapshot_ids?: string[];
|
||||
input_snapshot_id?: string | null;
|
||||
research_id?: string | null;
|
||||
parent_run_id?: string | null;
|
||||
hypothesis?: string | null;
|
||||
|
||||
@@ -18,6 +18,13 @@ import type { WorkspacePage } from "../ai/workspace";
|
||||
const navigation = [
|
||||
{ id: "home", label: "首页看板", icon: IconGridView, 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: "templates",
|
||||
|
||||
@@ -93,6 +93,14 @@ export function JobPanel({
|
||||
{job.checkpoint.dates_completed} / {job.checkpoint.dates_total} 天
|
||||
</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" && (
|
||||
<p className="muted">
|
||||
{job.total === 0
|
||||
|
||||
+378
-1324
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>
|
||||
);
|
||||
}
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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%;
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
});
|
||||
@@ -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}
|
||||
</Button>
|
||||
))}
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
import {
|
||||
ResearchDataInput,
|
||||
researchSelection,
|
||||
} from "../preparations/ResearchDataInput";
|
||||
import { ResearchList, ResearchDetail } from "./ResearchList";
|
||||
import { useEffect, useRef, useState } from "react";
|
||||
import {
|
||||
@@ -66,9 +70,7 @@ export function FeaturesPage({
|
||||
`/research/assets?kind=feature&q=${encodeURIComponent(q)}&offset=${(page - 1) * 25}`,
|
||||
{ signal: controller.signal },
|
||||
),
|
||||
api<{ items: InputSnapshot[] }>("/research/inputs", {
|
||||
signal: controller.signal,
|
||||
}),
|
||||
Promise.resolve({ items: inputs }),
|
||||
])
|
||||
.then(([list, fixed]) => {
|
||||
setItems(list.items);
|
||||
@@ -128,7 +130,7 @@ export function FeaturesPage({
|
||||
body: JSON.stringify({
|
||||
kind: "feature",
|
||||
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))
|
||||
.map((i) => ({
|
||||
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}
|
||||
busy={busy === "generate"}
|
||||
@@ -330,7 +332,7 @@ export function FeaturesPage({
|
||||
await post<FeatureAsset>("/research/generate", {
|
||||
name: draft.name,
|
||||
hypothesis: draft.hypothesis,
|
||||
input_ids: draft.input_ids,
|
||||
...researchSelection(draft.input_ids, inputs),
|
||||
method: "feature",
|
||||
}),
|
||||
);
|
||||
@@ -347,23 +349,21 @@ export function FeaturesPage({
|
||||
</section>
|
||||
<label>
|
||||
固定输入
|
||||
<ResearchSelect
|
||||
label="特征固定输入"
|
||||
multiple
|
||||
filter
|
||||
value={draft.input_ids}
|
||||
optionList={inputs.map((i) => ({
|
||||
value: i.id,
|
||||
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[] })}
|
||||
<ResearchDataInput
|
||||
label="选择特征数据准备"
|
||||
ids={draft.input_ids}
|
||||
inputs={inputs}
|
||||
onChange={(ids, rows) => {
|
||||
setInputs(rows);
|
||||
setDraft({ ...draft, input_ids: ids });
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
{inputs
|
||||
.filter((i) => draft.input_ids.includes(i.id))
|
||||
.map((i) => (
|
||||
<details key={i.id}>
|
||||
<summary>{i.dataset_id} 的固定字段</summary>
|
||||
<summary>{i.name} 的固定字段</summary>
|
||||
<p className="muted">{i.field_ids.join("、")}</p>
|
||||
</details>
|
||||
))}
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
import {
|
||||
ResearchDataInput,
|
||||
researchSelection,
|
||||
} from "../preparations/ResearchDataInput";
|
||||
import { useEffect, useState } from "react";
|
||||
import {
|
||||
Banner,
|
||||
@@ -54,7 +58,7 @@ export function FlowLaunchForm({
|
||||
useEffect(() => {
|
||||
const c = new AbortController();
|
||||
Promise.all([
|
||||
api<{ items: InputSnapshot[] }>("/research/inputs", { signal: c.signal }),
|
||||
Promise.resolve({ items: inputs }),
|
||||
api<{ items: Asset[] }>("/research/assets?kind=template&limit=100", {
|
||||
signal: c.signal,
|
||||
}),
|
||||
@@ -73,7 +77,7 @@ export function FlowLaunchForm({
|
||||
setConfirmation({
|
||||
request_id: crypto.randomUUID(),
|
||||
name,
|
||||
input_ids: ids,
|
||||
...researchSelection(ids, inputs),
|
||||
hypothesis,
|
||||
settings,
|
||||
budget,
|
||||
@@ -133,22 +137,17 @@ export function FlowLaunchForm({
|
||||
</label>
|
||||
<label>
|
||||
固定数据范围
|
||||
<ResearchSelect
|
||||
label="自动研究固定输入"
|
||||
multiple
|
||||
filter
|
||||
value={ids}
|
||||
optionList={inputs.map((i) => ({
|
||||
value: i.id,
|
||||
label: `${i.dataset_id} · ${i.scope.region}/${i.scope.universe}/D${i.scope.delay} · ${i.id.slice(0, 8)}`,
|
||||
}))}
|
||||
onChange={(v) => {
|
||||
const next = v as string[];
|
||||
setIds(next);
|
||||
const first = inputs.find((i) => i.id === next[0]);
|
||||
<ResearchDataInput
|
||||
label="选择研究数据准备"
|
||||
ids={ids}
|
||||
inputs={inputs}
|
||||
onChange={(ids, rows) => {
|
||||
setInputs(rows);
|
||||
setIds(ids);
|
||||
const first = rows.find((r) => r.id === ids[0]);
|
||||
if (first)
|
||||
setSettings((s) => ({
|
||||
...s,
|
||||
setSettings((old) => ({
|
||||
...old,
|
||||
region: first.scope.region,
|
||||
universe: first.scope.universe,
|
||||
delay: first.scope.delay,
|
||||
@@ -290,8 +289,10 @@ export function FlowLaunchForm({
|
||||
</p>
|
||||
)}
|
||||
<p>
|
||||
{confirmation.input_ids.length} 个固定输入 · {settings.region}/
|
||||
{settings.universe}/D{settings.delay}
|
||||
{confirmation.input_ids.length +
|
||||
(confirmation.preparation_refs?.length ?? 0)}{" "}
|
||||
个固定输入 · {settings.region}/{settings.universe}/D
|
||||
{settings.delay}
|
||||
</p>
|
||||
<ul>
|
||||
{confirmation.input_ids.map((id) => {
|
||||
@@ -299,7 +300,7 @@ export function FlowLaunchForm({
|
||||
return (
|
||||
<li key={id}>
|
||||
{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.slice(0, 8)}
|
||||
</li>
|
||||
@@ -315,7 +316,9 @@ export function FlowLaunchForm({
|
||||
{confirmation.template_id && (
|
||||
<p>
|
||||
初始模板:
|
||||
{templates.find((t) => t.id === confirmation.template_id)?.name}{" "}
|
||||
{
|
||||
templates.find((t) => t.id === confirmation.template_id)?.name
|
||||
}{" "}
|
||||
· v{confirmation.template_version}
|
||||
</p>
|
||||
)}
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
import {
|
||||
ResearchDataInput,
|
||||
researchSelection,
|
||||
} from "../preparations/ResearchDataInput";
|
||||
import { ResearchList, ResearchDetail } from "./ResearchList";
|
||||
import { ResearchSelect } from "./ResearchSelect";
|
||||
import { useEffect, useRef, useState } from "react";
|
||||
@@ -101,7 +105,7 @@ export function ResearchWorkspace({
|
||||
api<{ items: Asset[]; total: number }>(
|
||||
`/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 }>(
|
||||
`/research/experiments?limit=25&offset=${historyPage * 25}`,
|
||||
),
|
||||
@@ -129,7 +133,11 @@ export function ResearchWorkspace({
|
||||
research_asset_id: asset?.id,
|
||||
research_experiment_id: experiment?.id,
|
||||
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]);
|
||||
useEffect(() => {
|
||||
@@ -180,17 +188,6 @@ export function ResearchWorkspace({
|
||||
setMethod("structure");
|
||||
}
|
||||
}, [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() {
|
||||
const data = await post<{
|
||||
content: {
|
||||
@@ -236,7 +233,7 @@ export function ResearchWorkspace({
|
||||
const next = await post<Asset>("/research/generate", {
|
||||
name: template.name,
|
||||
hypothesis,
|
||||
input_ids: inputIds,
|
||||
...researchSelection(inputIds, inputs),
|
||||
parent_alpha_ids: parent ? parent.split(/[,,\s]+/).filter(Boolean) : [],
|
||||
method: page === "variants" ? "structure" : "template",
|
||||
});
|
||||
@@ -274,7 +271,7 @@ export function ResearchWorkspace({
|
||||
const next = await post<Experiment>("/research/experiments", {
|
||||
asset_id: asset!.id,
|
||||
version: asset!.version,
|
||||
input_ids: inputIds,
|
||||
...researchSelection(inputIds, inputs),
|
||||
hypothesis,
|
||||
settings,
|
||||
mode,
|
||||
@@ -580,20 +577,26 @@ export function ResearchWorkspace({
|
||||
)}
|
||||
<label>
|
||||
固定研究输入
|
||||
<ResearchSelect
|
||||
multiple
|
||||
filter
|
||||
label="固定研究输入"
|
||||
value={inputIds}
|
||||
optionList={inputs.map((input) => ({
|
||||
value: input.id,
|
||||
label: `${input.dataset_id} · ${input.scope.region}/${input.scope.universe}/D${input.scope.delay} · ${input.field_ids.length} 字段 · ${input.id.slice(0, 8)}`,
|
||||
}))}
|
||||
onChange={(value) => selectInputs(value as string[])}
|
||||
<ResearchDataInput
|
||||
label="选择数据准备"
|
||||
ids={inputIds}
|
||||
inputs={inputs}
|
||||
onChange={(ids, rows) => {
|
||||
setInputs(rows);
|
||||
setInputIds(ids);
|
||||
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,
|
||||
}));
|
||||
}}
|
||||
/>
|
||||
</label>
|
||||
<p className="research-hint">
|
||||
在数据目录中保存字段选择。跨数据集分别关联输入,跨市场使用目标范围的独立输入。
|
||||
从数据准备选择集合;每个集合可包含同范围下多个数据集的字段。
|
||||
</p>
|
||||
{page === "variants" && method === "settings" ? (
|
||||
<label>
|
||||
@@ -624,7 +627,7 @@ export function ResearchWorkspace({
|
||||
references={[
|
||||
...selectedInputs.map((input) => ({
|
||||
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
|
||||
? [
|
||||
@@ -857,7 +860,7 @@ export function ResearchWorkspace({
|
||||
setExperiment(
|
||||
await post("/research/variants/settings", {
|
||||
alpha_id: parent,
|
||||
input_ids: inputIds,
|
||||
...researchSelection(inputIds, inputs),
|
||||
...(hypothesis ? { hypothesis } : {}),
|
||||
}),
|
||||
);
|
||||
|
||||
@@ -62,19 +62,28 @@ export function SourceDetails({
|
||||
打开研究会话
|
||||
</Button>
|
||||
)}
|
||||
{source.template_input_id && (
|
||||
{[
|
||||
...new Set([
|
||||
...(source.input_snapshot_ids ?? []),
|
||||
...(source.input_snapshot_id ? [source.input_snapshot_id] : []),
|
||||
]),
|
||||
].map((id, index) => (
|
||||
<Button
|
||||
key={id}
|
||||
onClick={() =>
|
||||
onAction({
|
||||
type: "open_research_input",
|
||||
input_id: source.template_input_id!,
|
||||
input_id: id,
|
||||
nonce: Date.now(),
|
||||
})
|
||||
}
|
||||
>
|
||||
查看研究输入
|
||||
{(source.input_snapshot_ids?.length ?? 0) > 1
|
||||
? ` ${index + 1}`
|
||||
: ""}
|
||||
</Button>
|
||||
)}
|
||||
))}
|
||||
{source.parent_run_id && (
|
||||
<Button
|
||||
onClick={() =>
|
||||
|
||||
@@ -21,6 +21,7 @@ export type FlowLaunch = {
|
||||
request_id: string;
|
||||
name: string;
|
||||
input_ids: string[];
|
||||
preparation_refs?: { id: string; version: number }[];
|
||||
hypothesis: string;
|
||||
settings: SimulationSettings;
|
||||
budget: Budget;
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import type { FieldRecord, PreparationRef } from "../preparations/types";
|
||||
import type { Candidate } from "../backtests/types";
|
||||
export type Variable = {
|
||||
kind:
|
||||
@@ -43,10 +44,14 @@ export type InputSnapshot = {
|
||||
universe: string;
|
||||
delay: 0 | 1;
|
||||
};
|
||||
dataset_id: string;
|
||||
name: string;
|
||||
dataset_ids: string[];
|
||||
fields: FieldRecord[];
|
||||
preparation_ref?: PreparationRef;
|
||||
field_ids: string[];
|
||||
field_types: Record<string, string>;
|
||||
collection_version: string;
|
||||
preparation_id?: string;
|
||||
preparation_version?: number;
|
||||
created_at: string;
|
||||
};
|
||||
export type ResearchCandidate = Candidate & {
|
||||
|
||||
@@ -139,6 +139,9 @@ export type Job = {
|
||||
};
|
||||
checkpoint: {
|
||||
date?: string;
|
||||
dataset_id?: string;
|
||||
datasets_completed?: number;
|
||||
datasets_total?: number;
|
||||
dates_completed?: number;
|
||||
dates_total?: number;
|
||||
alpha_id?: string;
|
||||
|
||||
+95
-238
@@ -1,14 +1,12 @@
|
||||
import { test, expect, type Page } from "@playwright/test";
|
||||
import { choosePreparation } from "./preparation-helpers";
|
||||
const headers = { "X-WQ-Request": "1" };
|
||||
const scope = {
|
||||
instrument_type: "EQUITY",
|
||||
region: "USA",
|
||||
universe: "TOP3000",
|
||||
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) {
|
||||
await page.goto("/#datasets");
|
||||
await page.getByLabel("密码", { exact: true }).fill("browser-test-password");
|
||||
@@ -20,273 +18,132 @@ async function setup(page: Page) {
|
||||
headers,
|
||||
data: { email: "test@example.com", password: "synthetic-only" },
|
||||
});
|
||||
const connect = await (
|
||||
const job = await (
|
||||
await page.request.post("/api/v1/account/connect", { headers })
|
||||
).json();
|
||||
await expect
|
||||
.poll(
|
||||
async () =>
|
||||
(
|
||||
await (
|
||||
await page.request.get(`/api/v1/sync-jobs/${connect.id}`)
|
||||
).json()
|
||||
).status,
|
||||
(await (await page.request.get(`/api/v1/sync-jobs/${job.id}`)).json())
|
||||
.status,
|
||||
)
|
||||
.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/catalog/datasets?${query}`)
|
||||
).json()
|
||||
).total,
|
||||
(await (await page.request.get(`/api/v1/sync-jobs/${job.id}`)).json())
|
||||
.status,
|
||||
)
|
||||
.toBe(3);
|
||||
await page.keyboard.press("Escape");
|
||||
await expect(
|
||||
page.getByRole("button", { name: "TEST 财务报表", exact: true }),
|
||||
).toBeVisible();
|
||||
.toBe("completed");
|
||||
}
|
||||
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();
|
||||
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
|
||||
.poll(
|
||||
async () =>
|
||||
(
|
||||
await (
|
||||
await page.request.get(
|
||||
`/api/v1/catalog/datasets/TEST_FIN/fields?${query}`,
|
||||
"/api/v1/catalog/datasets/TEST_FIN?region=USA&universe=TOP3000&delay=1",
|
||||
)
|
||||
).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 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(
|
||||
page.getByRole("button", { name: "已保存输入 (2)" }),
|
||||
page.locator("p").filter({ hasText: /浏览器准备集合 ·.*122 字段/ }),
|
||||
).toBeVisible();
|
||||
expect(errors).toEqual([]);
|
||||
});
|
||||
|
||||
test("窄屏抽屉、键盘隔离和研究范围切换", async ({ page }) => {
|
||||
await page.setViewportSize({ width: 390, height: 844 });
|
||||
test("在线字段跨页多选直接准备,本地只展示完整同步字段", async ({ 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
|
||||
.getByRole("button", { name: "TEST 字段 000", exact: true })
|
||||
.getByRole("navigation", { name: "主导航" })
|
||||
.getByRole("button", { name: "字段目录", exact: true })
|
||||
.click();
|
||||
await expect(
|
||||
page.getByRole("dialog", { name: "字段详情", exact: true }),
|
||||
).toBeVisible();
|
||||
await page.keyboard.press("Escape");
|
||||
await expect(fields).toBeVisible();
|
||||
expect(
|
||||
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 })
|
||||
const online = page.getByRole("tabpanel", { name: "worldquant接口" });
|
||||
await expect(online.getByText("TEST_FIN_001", { exact: true })).toBeVisible();
|
||||
await online
|
||||
.getByRole("row")
|
||||
.filter({ hasText: "TEST_FIN_001" })
|
||||
.locator("span.semi-checkbox")
|
||||
.click();
|
||||
await page.getByLabel("研究备注", { exact: true }).fill("未保存草稿需要恢复");
|
||||
await page
|
||||
.getByRole("dialog", { name: "字段详情", exact: true })
|
||||
.getByRole("button", { name: "AI 助手", exact: true })
|
||||
await online.locator(".semi-page-item").filter({ hasText: /^2$/ }).click();
|
||||
await online
|
||||
.getByRole("row")
|
||||
.filter({ hasText: "TEST_FIN_025" })
|
||||
.locator("span.semi-checkbox")
|
||||
.click();
|
||||
await expect(fields).not.toBeVisible();
|
||||
await page.keyboard.press("Escape");
|
||||
await expect(page.getByLabel("研究备注", { exact: true })).toHaveValue(
|
||||
"未保存草稿需要恢复",
|
||||
);
|
||||
for (const width of [850, 390, 1920]) {
|
||||
await page.setViewportSize({ width, height: 900 });
|
||||
await expect(
|
||||
page.getByRole("button", { name: "关闭详情", exact: true }),
|
||||
).toBeVisible();
|
||||
expect(
|
||||
await page.evaluate(
|
||||
() => document.documentElement.scrollWidth <= innerWidth,
|
||||
),
|
||||
).toBe(true);
|
||||
}
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.mouse.click(10, 400);
|
||||
await expect(
|
||||
page.getByRole("dialog", { name: "字段详情", exact: true }),
|
||||
).not.toBeVisible();
|
||||
await expect(fields).toBeVisible();
|
||||
await expect(fields).toContainText("122 / 123 个字段");
|
||||
await page.mouse.click(10, 400);
|
||||
await expect(fields).not.toBeVisible();
|
||||
expect(errors).toEqual([]);
|
||||
await expect(online.getByText(/已选 2 个字段/)).toBeVisible();
|
||||
await online
|
||||
.getByRole("button", { name: "加入数据准备", exact: true })
|
||||
.click();
|
||||
const add = page.getByRole("dialog", { name: "加入数据准备", exact: true });
|
||||
await add.getByLabel("新集合名称").fill("在线选取集合");
|
||||
await add.getByRole("button", { name: "新建集合并加入" }).click();
|
||||
await expect(add).not.toBeVisible();
|
||||
const collections = await (
|
||||
await page.request.get("/api/v1/data-preparations?q=在线选取集合")
|
||||
).json();
|
||||
expect(collections.items[0].field_count).toBe(2);
|
||||
await page.getByRole("tab", { name: "本地同步", exact: true }).click();
|
||||
const local = page.getByRole("tabpanel", { name: "本地同步" });
|
||||
// Other browser tests share this isolated server, so verify the unsynced dataset explicitly.
|
||||
await local.getByLabel("字段数据集").fill("TEST_NEWS");
|
||||
await local.getByRole("button", { name: "查询", exact: true }).click();
|
||||
await expect(local.getByText("共 0 个字段", { exact: true })).toBeVisible();
|
||||
await page.screenshot({
|
||||
path: "test-results/preparations-fields.png",
|
||||
fullPage: true,
|
||||
});
|
||||
|
||||
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();
|
||||
});
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
import { choosePreparation } from "./preparation-helpers";
|
||||
import { expect, test } from "@playwright/test";
|
||||
const headers = { "X-WQ-Request": "1" };
|
||||
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}`)
|
||||
).json();
|
||||
const fixed = await (
|
||||
await page.request.post("/api/v1/catalog/inputs", {
|
||||
await page.request.post("/api/v1/data-preparations/from-dataset", {
|
||||
headers,
|
||||
data: {
|
||||
scope,
|
||||
dataset_id: "TEST_FIN",
|
||||
collection_version: fields.collection_version,
|
||||
selection: "all",
|
||||
},
|
||||
})
|
||||
).json();
|
||||
@@ -175,11 +175,7 @@ test("native QuantFlow canvas saves versions, branches and restores artifacts",
|
||||
.getByRole("button", { name: "启动已保存流程", exact: true })
|
||||
.click();
|
||||
await page.getByLabel("自动研究假设").fill("画布条件分支与固定版本");
|
||||
await page.getByRole("combobox", { name: "自动研究固定输入" }).click();
|
||||
await page
|
||||
.getByRole("option")
|
||||
.filter({ hasText: fixed.id.slice(0, 8) })
|
||||
.click();
|
||||
await choosePreparation(page, fixed.name, "选择研究数据准备");
|
||||
await page.getByLabel("自动研究假设").click();
|
||||
for (const [label, value] of [
|
||||
["最大轮数", "1"],
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import { choosePreparation } from "./preparation-helpers";
|
||||
import { expect, test } from "@playwright/test";
|
||||
const headers = { "X-WQ-Request": "1" };
|
||||
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}`)
|
||||
).json();
|
||||
const fixed = await (
|
||||
await page.request.post("/api/v1/catalog/inputs", {
|
||||
await page.request.post("/api/v1/data-preparations/from-dataset", {
|
||||
headers,
|
||||
data: {
|
||||
scope,
|
||||
dataset_id: "TEST_FIN",
|
||||
collection_version: fields.collection_version,
|
||||
selection: "all",
|
||||
},
|
||||
})
|
||||
).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.getByLabel("自动研究名称").fill("浏览器有限研究");
|
||||
await page.getByLabel("自动研究假设").fill("研究排名稳定性");
|
||||
await page.getByRole("combobox", { name: "自动研究固定输入" }).click();
|
||||
await page
|
||||
.getByRole("option")
|
||||
.filter({ hasText: fixed.id.slice(0, 8) })
|
||||
.click();
|
||||
await choosePreparation(page, fixed.name, "选择研究数据准备");
|
||||
await page.getByLabel("自动研究名称").click();
|
||||
for (const [label, value] of [
|
||||
["最大轮数", "1"],
|
||||
|
||||
@@ -68,17 +68,27 @@ test("fixed input, AI approval, backtest and research source remain connected",
|
||||
const fields = await (
|
||||
await page.request.get(`/api/v1/catalog/datasets/TEST_FIN/fields?${params}`)
|
||||
).json();
|
||||
const saved = await page.request.post("/api/v1/catalog/inputs", {
|
||||
const saved = await page.request.post(
|
||||
"/api/v1/data-preparations/from-dataset",
|
||||
{
|
||||
headers,
|
||||
data: {
|
||||
scope,
|
||||
dataset_id: "TEST_FIN",
|
||||
collection_version: fields.collection_version,
|
||||
selection: "all",
|
||||
},
|
||||
});
|
||||
},
|
||||
);
|
||||
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 (
|
||||
await page.request.post("/api/v1/ai/conversations", { headers })
|
||||
).json();
|
||||
@@ -96,7 +106,7 @@ test("fixed input, AI approval, backtest and research source remain connected",
|
||||
page: "datasets",
|
||||
catalog_scope: scope,
|
||||
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();
|
||||
await expect(chat).not.toBeVisible();
|
||||
const inputSheet = page.getByRole("dialog", {
|
||||
name: "模板输入草稿",
|
||||
name: "研究输入快照",
|
||||
exact: true,
|
||||
});
|
||||
await expect(
|
||||
inputSheet.getByRole("heading", { name: "输入草稿已保存", exact: true }),
|
||||
).toBeVisible();
|
||||
await inputSheet
|
||||
.getByRole("button", { name: "用此输入研究", exact: true })
|
||||
.click();
|
||||
await expect(
|
||||
chat.locator(".research-ai-reference").filter({ hasText: "固定输入" }),
|
||||
).toContainText(input.id);
|
||||
await expect(inputSheet).toContainText("TEST_FIN_001");
|
||||
await expect(inputSheet).toContainText("TEST_FIN");
|
||||
expect(errors).toEqual([]);
|
||||
});
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import { choosePreparation } from "./preparation-helpers";
|
||||
import { expect, test } from "@playwright/test";
|
||||
const headers = { "X-WQ-Request": "1" };
|
||||
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}`)
|
||||
).json();
|
||||
const fixed = await (
|
||||
await page.request.post("/api/v1/catalog/inputs", {
|
||||
await page.request.post("/api/v1/data-preparations/from-dataset", {
|
||||
headers,
|
||||
data: {
|
||||
scope,
|
||||
dataset_id: "TEST_FIN",
|
||||
collection_version: fields.collection_version,
|
||||
selection: "all",
|
||||
},
|
||||
})
|
||||
).json();
|
||||
@@ -67,11 +67,7 @@ test("feature conversion, saved views and immutable evaluations", async ({
|
||||
await page
|
||||
.getByLabel("特征经济假设", { exact: true })
|
||||
.fill("排名可降低异常值影响");
|
||||
await page.getByRole("combobox", { name: "特征固定输入" }).click();
|
||||
await page
|
||||
.getByRole("option")
|
||||
.filter({ hasText: fixed.id.slice(0, 8) })
|
||||
.click();
|
||||
await choosePreparation(page, fixed.name, "选择特征数据准备");
|
||||
await page.getByLabel("特征方案名称", { exact: true }).click();
|
||||
await page.getByRole("button", { name: "添加处理步骤" }).click();
|
||||
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=特征转换模板",
|
||||
)
|
||||
).json();
|
||||
expect(templates.items[0].provenance.feature.content.input_ids).toEqual([
|
||||
fixed.id,
|
||||
]);
|
||||
const snapshot = await (
|
||||
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.getByRole("button", { name: "Alpha 管理", exact: true }).click();
|
||||
await page.getByRole("button", { name: "保存为新视图" }).click();
|
||||
|
||||
@@ -1,9 +1,17 @@
|
||||
import { choosePreparation } from "./preparation-helpers";
|
||||
import { expect, test } from "@playwright/test";
|
||||
|
||||
// UI contract fixtures: no external model call or simulation is started by these tests.
|
||||
const fixed = {
|
||||
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_types: { TEST_FIN_001: "MATRIX" },
|
||||
collection_version: "fixed-v1",
|
||||
@@ -20,8 +28,19 @@ for (const mode of ["templates", "variants", "features"] as const) {
|
||||
}) => {
|
||||
const errors: string[] = [];
|
||||
page.on("pageerror", (error) => errors.push(error.message));
|
||||
await page.route("**/api/v1/research/inputs", (route) =>
|
||||
route.fulfill({ json: { items: [fixed] } }),
|
||||
await page.route("**/api/v1/data-preparations?*", (route) =>
|
||||
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(
|
||||
"**/api/v1/research/assets/semi-generated/versions",
|
||||
@@ -49,23 +68,18 @@ for (const mode of ["templates", "variants", "features"] as const) {
|
||||
exact: true,
|
||||
});
|
||||
await expect(generate).toBeDisabled();
|
||||
await sheet
|
||||
.getByRole("combobox", {
|
||||
name: feature ? "特征固定输入" : "固定研究输入",
|
||||
exact: true,
|
||||
})
|
||||
.click();
|
||||
await page
|
||||
.getByRole("option")
|
||||
.filter({ hasText: fixed.id.slice(0, 8) })
|
||||
.click();
|
||||
await choosePreparation(
|
||||
page,
|
||||
fixed.name,
|
||||
feature ? "选择特征数据准备" : "选择数据准备",
|
||||
);
|
||||
await input.click();
|
||||
if (mode === "variants")
|
||||
await sheet.getByLabel("种子 Alpha", { exact: true }).fill("TEST0001");
|
||||
const hypothesis = "检验 <收益> 与规模\n第二轮保留原始假设";
|
||||
await input.fill(hypothesis);
|
||||
await expect(
|
||||
sheet.locator(".research-ai-reference").filter({ hasText: fixed.id }),
|
||||
sheet.locator("p").filter({ hasText: /Semi 准备集合 ·/ }),
|
||||
).toBeVisible();
|
||||
let release!: () => void;
|
||||
const pending = new Promise<void>((resolve) => {
|
||||
@@ -117,7 +131,8 @@ for (const mode of ["templates", "variants", "features"] as const) {
|
||||
.poll(() => requestBody)
|
||||
.toMatchObject({
|
||||
hypothesis,
|
||||
input_ids: [fixed.id],
|
||||
input_ids: [],
|
||||
preparation_refs: [fixed.preparation_ref],
|
||||
method: feature
|
||||
? "feature"
|
||||
: mode === "variants"
|
||||
|
||||
Reference in New Issue
Block a user