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

This commit is contained in:
yuxuanhui
2026-09-12 01:24:02 +08:00
parent 849f86fef7
commit 394438e753
82 changed files with 4146 additions and 2076 deletions
+1
View File
@@ -0,0 +1 @@
"""Data preparation collections and immutable research inputs."""
+80
View File
@@ -0,0 +1,80 @@
"""Shared collection and field-query contracts."""
from datetime import datetime
from typing import Literal
from pydantic import Field, model_validator
from ..catalog.contracts import Scope
from ..schemas import Contract
class FieldFilters(Scope):
q: str = Field(default="", max_length=300)
dataset_id: str | None = None
category: str | None = None
subcategory: str | None = None
field_type: str | None = None
coverage_min: float | None = Field(default=None, ge=0, le=1)
coverage_max: float | None = Field(default=None, ge=0, le=1)
user_count_min: int | None = Field(default=None, ge=0)
user_count_max: int | None = Field(default=None, ge=0)
alpha_count_min: int | None = Field(default=None, ge=0)
alpha_count_max: int | None = Field(default=None, ge=0)
synced_from: datetime | None = None
synced_to: datetime | None = None
sort: Literal["id", "name", "dataset_id", "coverage", "user_count", "alpha_count", "synced_at"] = "id"
direction: Literal["asc", "desc"] = "asc"
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
@model_validator(mode="after")
def ranges(self):
for key in ("coverage", "user_count", "alpha_count"):
low, high = getattr(self, key + "_min"), getattr(self, key + "_max")
if low is not None and high is not None and low > high:
raise ValueError("筛选下限不能超过上限")
return self
class FieldReference(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
field_id: str = Field(min_length=1, max_length=200)
source: Literal["local", "worldquant"] = "local"
collection_version: str | None = None
class PreparationCreate(Contract):
name: str = Field(min_length=1, max_length=200)
note: str = Field(default="", max_length=20000)
scope: Scope
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
class PreparationVersion(Contract):
version: int = Field(ge=1)
class PreparationEdit(PreparationVersion):
name: str = Field(min_length=1, max_length=200)
note: str = Field(default="", max_length=20000)
class MemberChange(PreparationVersion):
fields: list[FieldReference] = Field(default_factory=list, max_length=10000)
remove_ids: list[str] = Field(default_factory=list, max_length=10000)
class PreparationReference(PreparationVersion):
id: str = Field(min_length=1, max_length=36)
class PreparationReferences(Contract):
items: list[PreparationReference] = Field(min_length=1, max_length=100)
class DatasetCopy(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
collection_version: str
+195
View File
@@ -0,0 +1,195 @@
"""Authenticated preparation and field-directory endpoints."""
from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import delete, select
from ..catalog.contracts import Scope
from ..models import CatalogScope, PreparationField, now
from ..security import require_auth
from .contracts import (
DatasetCopy,
FieldFilters,
MemberChange,
PreparationCreate,
PreparationEdit,
PreparationReferences,
PreparationVersion,
)
from .service import Preparations
router = APIRouter(prefix="/api/v1", dependencies=[Depends(require_auth)], tags=["data-preparations"])
@router.get("/catalog/local-scopes")
async def local_scopes(request: Request):
async with request.app.state.sessions() as db:
rows = await db.scalars(select(CatalogScope))
options = {}
for row in rows:
key = (row.scope["instrument_type"], row.scope["region"], row.scope["delay"])
option = options.setdefault(
key, {k: row.scope[k] for k in ("instrument_type", "region", "delay")}
)
option.setdefault("universes", []).append(row.scope["universe"])
return {"instrument_options": list(options.values())}
@router.get("/catalog/fields")
async def local_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
async with request.app.state.sessions() as db:
return await Preparations(db).fields(filters)
@router.get("/catalog/worldquant/fields")
async def online_fields(request: Request, filters: Annotated[FieldFilters, Query()]):
async with request.app.state.sessions() as db:
return await Preparations(db, request.app.state.runner.client).online_fields(filters)
@router.get("/data-preparations")
async def preparations(
request: Request,
q: str = "",
scope_key: str | None = None,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
):
async with request.app.state.sessions() as db:
return await Preparations(db).list(q, scope_key, limit, offset)
@router.post("/data-preparations", status_code=201)
async def create(request: Request, body: PreparationCreate):
async with request.app.state.sessions.begin() as db:
service = Preparations(db, request.app.state.runner.client)
fields = await service.resolve_fields(body.scope, body.fields)
return await service.create(body.name, body.note, body.scope, fields)
@router.post("/data-preparations/from-dataset", status_code=201)
async def from_dataset(request: Request, body: DatasetCopy):
async with request.app.state.sessions.begin() as db:
return await Preparations(db).copy_dataset(body)
@router.post("/data-preparations/batch-delete")
async def batch_delete(request: Request, body: PreparationReferences):
async with request.app.state.sessions.begin() as db:
return await Preparations(db).remove(body.items)
@router.post("/data-preparations/freeze", status_code=201)
async def freeze(request: Request, body: PreparationReferences):
async with request.app.state.sessions.begin() as db:
return {"items": await Preparations(db).freeze(body.items)}
@router.get("/research/input-snapshots/{snapshot_id}")
async def snapshot(request: Request, snapshot_id: str):
async with request.app.state.sessions() as db:
return await Preparations(db).snapshot(snapshot_id)
@router.get("/data-preparations/{preparation_id}")
async def detail(request: Request, preparation_id: str):
async with request.app.state.sessions() as db:
service = Preparations(db)
return await service.output(await service.get(preparation_id))
@router.patch("/data-preparations/{preparation_id}")
async def edit(request: Request, preparation_id: str, body: PreparationEdit):
async with request.app.state.sessions.begin() as db:
service = Preparations(db)
row = await service.get(preparation_id, body.version, lock=True)
row.name, row.note, row.updated_at, row.version = body.name, body.note, now(), row.version + 1
return await service.output(row)
@router.delete("/data-preparations/{preparation_id}")
async def remove(request: Request, preparation_id: str, version: int = Query(ge=1)):
from .contracts import PreparationReference
async with request.app.state.sessions.begin() as db:
return await Preparations(db).remove([PreparationReference(id=preparation_id, version=version)])
@router.post("/data-preparations/{preparation_id}/copy", status_code=201)
async def copy(request: Request, preparation_id: str, body: PreparationVersion):
async with request.app.state.sessions.begin() as db:
service = Preparations(db)
row = await service.get(preparation_id, body.version, lock=True)
fields = [
f.content
for f in await db.scalars(
select(PreparationField).where(PreparationField.preparation_id == row.id)
)
]
return await service.create(
(row.name + " 副本")[:200], row.note, Scope.model_validate(row.scope), fields
)
@router.get("/data-preparations/{preparation_id}/fields")
async def members(
request: Request,
preparation_id: str,
q: str = "",
dataset_id: str | None = None,
limit: int = Query(25, ge=1, le=100),
offset: int = Query(0, ge=0),
):
async with request.app.state.sessions() as db:
return await Preparations(db).members(preparation_id, q, dataset_id, limit, offset)
@router.patch("/data-preparations/{preparation_id}/fields")
async def change_members(request: Request, preparation_id: str, body: MemberChange):
async with request.app.state.sessions.begin() as db:
service = Preparations(db, request.app.state.runner.client)
row = await service.get(preparation_id, body.version, lock=True)
fields = await service.resolve_fields(Scope.model_validate(row.scope), body.fields)
await service.add(row, fields)
if body.remove_ids:
present = set(
await db.scalars(
select(PreparationField.field_id).where(
PreparationField.preparation_id == row.id,
PreparationField.field_id.in_(body.remove_ids),
)
)
)
if set(body.remove_ids) - present:
raise HTTPException(422, "移除项含不属于该集合的字段")
await db.execute(
delete(PreparationField).where(
PreparationField.preparation_id == row.id, PreparationField.field_id.in_(body.remove_ids)
)
)
row.version, row.updated_at = row.version + 1, now()
return await service.output(row)
@router.get("/data-preparations/{preparation_id}/selection")
async def selection(request: Request, preparation_id: str, version: int = Query(ge=1)):
async with request.app.state.sessions() as db:
service = Preparations(db)
row = await service.get(preparation_id, version)
fields = [
f.content
for f in await db.scalars(
select(PreparationField)
.where(PreparationField.preparation_id == row.id)
.order_by(PreparationField.field_id)
)
]
return {
**await service.output(row),
"fields": fields,
"field_ids": [f["id"] for f in fields],
"field_types": {f["id"]: f["field_type"] for f in fields},
"dataset_ids": sorted({f["dataset_id"] for f in fields}),
"preparation_ref": {"id": row.id, "version": row.version},
}
+450
View File
@@ -0,0 +1,450 @@
"""Collection operations own validation; callers own transactions and authorization."""
from uuid import uuid4
from fastapi import HTTPException
from sqlalchemy import delete, func, or_, select
from sqlalchemy.orm import aliased
from ..catalog.contracts import EntryOutput, Scope
from ..catalog.research_metadata import upstream
from ..catalog.service import Catalog
from ..catalog.sync import identifier, label, normalize
from ..models import (
CatalogBatch,
CatalogDataset,
CatalogEntry,
CatalogScope,
DataPreparation,
PreparationField,
ResearchInputSnapshot,
now,
)
from ..research.serialization import encode_snapshot
from .contracts import FieldFilters
def page(items, total, limit, offset, **extra):
return dict(
items=items, total=total, limit=limit, offset=offset, has_more=offset + len(items) < total, **extra
)
def contains(value):
return "%" + value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
class Preparations:
def __init__(self, db, client=None):
self.db, self.client = db, client
async def fields(self, filters):
scope = await self.db.get(CatalogScope, filters.key())
owner = aliased(CatalogEntry)
query = (
select(
CatalogEntry,
CatalogBatch.dataset_id,
owner.name.label("dataset_name"),
owner.category,
owner.subcategory,
)
.join(CatalogBatch, CatalogEntry.batch_id == CatalogBatch.id)
.join(
CatalogDataset,
(CatalogDataset.field_version == CatalogBatch.id)
& (CatalogDataset.scope_key == filters.key()),
)
.outerjoin(
owner,
(owner.id == CatalogBatch.dataset_id)
& (owner.batch_id == (scope.catalog_version if scope else None)),
)
.where(CatalogBatch.complete.is_(True))
)
if filters.q:
query = query.where(
or_(
*[
column.ilike(contains(filters.q), escape="\\")
for column in (
CatalogEntry.id,
CatalogEntry.name,
CatalogEntry.description,
CatalogBatch.dataset_id,
owner.name,
)
]
)
)
for key in ("dataset_id", "field_type", "category", "subcategory"):
value = getattr(filters, key)
column = (
CatalogBatch.dataset_id
if key == "dataset_id"
else func.coalesce(getattr(CatalogEntry, key), getattr(owner, key))
if key in ("category", "subcategory")
else getattr(CatalogEntry, key)
)
if value:
query = query.where(column == value)
for key in ("coverage", "user_count", "alpha_count"):
low, high = getattr(filters, key + "_min"), getattr(filters, key + "_max")
if low is not None:
query = query.where(getattr(CatalogEntry, key) >= low)
if high is not None:
query = query.where(getattr(CatalogEntry, key) <= high)
if filters.synced_from:
query = query.where(CatalogEntry.synced_at >= filters.synced_from)
if filters.synced_to:
query = query.where(CatalogEntry.synced_at <= filters.synced_to)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
column = (
CatalogBatch.dataset_id if filters.sort == "dataset_id" else getattr(CatalogEntry, filters.sort)
)
query = query.order_by(
(column.desc() if filters.direction == "desc" else column.asc()).nulls_last(),
CatalogBatch.dataset_id,
CatalogEntry.id,
)
rows = (await self.db.execute(query.limit(filters.limit).offset(filters.offset))).all()
items = [
{
**self.local_field(row, dataset_id, name, filters),
"category": row.category or category,
"subcategory": row.subcategory or subcategory,
}
for row, dataset_id, name, category, subcategory in rows
]
return page(items, total, filters.limit, filters.offset)
@staticmethod
def local_field(row, dataset_id, name, scope):
return encode_snapshot(
dict(
**EntryOutput.model_validate(row, from_attributes=True).model_dump(
exclude={"scope", "dataset_id", "collection_version"}
),
field_id=row.id,
dataset_id=dataset_id,
dataset_name=name or dataset_id,
collection_version=row.batch_id,
scope=scope.model_dump(include=set(Scope.model_fields)),
source="local",
fetched_at=row.synced_at,
)
)
async def online_fields(self, filters):
if self.client is None:
raise HTTPException(409, "请先连接 WorldQuant")
params = dict(
instrumentType=filters.instrument_type,
region=filters.region,
universe=filters.universe,
delay=filters.delay,
limit=filters.limit,
offset=filters.offset,
)
for key, remote in (("q", "search"), ("dataset_id", "dataset.id"), ("field_type", "type")):
value = getattr(filters, key)
if value:
params[remote] = value
for key, remote in (
("coverage", "coverage"),
("user_count", "userCount"),
("alpha_count", "alphaCount"),
):
for suffix, op in (("min", ">"), ("max", "<")):
value = getattr(filters, key + "_" + suffix)
if value is not None:
params[remote + op] = value
raw = await upstream(self.client.get("/data-fields", params))
rows = raw.get("results")
if not isinstance(rows, list):
raise HTTPException(502, "平台字段列表格式无法识别")
items = []
for row in rows:
if not isinstance(row, dict):
raise HTTPException(502, "平台字段记录格式无法识别")
dataset = row.get("dataset")
owner = dataset.get("id") if isinstance(dataset, dict) else dataset
try:
owner = identifier(owner)
data = normalize(row, owner)
except Exception as exc:
from ..worldquant import WqError
if isinstance(exc, WqError):
raise HTTPException(502, str(exc)) from None
raise
if filters.dataset_id and owner != filters.dataset_id:
raise HTTPException(502, "平台返回其他数据集字段")
for remote_key, key in (
("instrumentType", "instrument_type"),
("instrument_type", "instrument_type"),
("region", "region"),
("universe", "universe"),
("delay", "delay"),
):
if remote_key in row and row[remote_key] != getattr(filters, key):
raise HTTPException(502, "平台字段范围与查询不一致")
items.append(
encode_snapshot(
dict(
**data,
field_id=data["id"],
dataset_id=owner,
dataset_name=label(dataset) or owner,
source="worldquant",
collection_version=None,
scope=filters.model_dump(include=set(Scope.model_fields)),
fetched_at=now(),
synced_at=None,
)
)
)
count = raw.get("count")
known_total = type(count) is int and count >= 0
more = (
bool(raw["next"])
if "next" in raw
else (filters.offset + len(items) < count if known_total else len(items) == filters.limit)
)
result = page(
items,
count if known_total else filters.offset + len(items) + int(more),
filters.limit,
filters.offset,
)
result.update(has_more=more, total_known=known_total)
return result
async def resolve_fields(self, scope, refs):
"""Resolve trusted source records before any collection member is written."""
if any(ref.scope.key() != scope.key() for ref in refs):
raise HTTPException(422, "不能跨区域、Top、Delay 或品种添加字段")
result = {}
for ref in refs:
if ref.source == "local":
dataset = await Catalog(self.db).dataset(scope, ref.dataset_id, lock=True)
if not dataset.field_version or dataset.field_version != ref.collection_version:
raise HTTPException(409, "字段来源已更新,请重新查询后添加")
entry = await self.db.get(CatalogEntry, (dataset.field_version, ref.field_id))
if not entry:
raise HTTPException(422, "字段不属于指定数据集")
scope_row = await self.db.get(CatalogScope, scope.key())
owner = await self.db.get(CatalogEntry, (scope_row.catalog_version, ref.dataset_id))
field = self.local_field(entry, ref.dataset_id, owner.name if owner else None, scope)
else:
offset, seen, field = 0, set(), None
while True:
response = await self.online_fields(
FieldFilters(
**scope.model_dump(),
q=ref.field_id,
dataset_id=ref.dataset_id,
limit=100,
offset=offset,
)
)
field = next((item for item in response["items"] if item["id"] == ref.field_id), None)
if field or not response["has_more"]:
break
ids = {item["id"] for item in response["items"]}
if not ids - seen:
raise HTTPException(502, "平台字段分页未前进")
seen.update(ids)
offset += len(response["items"])
if not field:
raise HTTPException(422, "在线字段已不可用,请重新查询")
if field["id"] in result and result[field["id"]]["dataset_id"] != field["dataset_id"]:
raise HTTPException(422, "同名字段的数据集归属冲突")
result[field["id"]] = field
return list(result.values())
async def get(self, preparation_id, version=None, lock=False):
query = select(DataPreparation).where(DataPreparation.id == preparation_id)
row = await self.db.scalar(query.with_for_update() if lock else query)
if not row:
raise HTTPException(404, "数据准备集合不存在")
if version is not None and row.version != version:
raise HTTPException(409, "集合已修改,请重新读取或选择;当前草稿已保留")
return row
async def output(self, row):
count, datasets = (
await self.db.execute(
select(func.count(), func.count(func.distinct(PreparationField.dataset_id))).where(
PreparationField.preparation_id == row.id
)
)
).one()
return encode_snapshot(
dict(
id=row.id,
name=row.name,
note=row.note,
scope=row.scope,
version=row.version,
field_count=count,
dataset_count=datasets,
created_at=row.created_at,
updated_at=row.updated_at,
)
)
async def list(self, q="", scope_key=None, limit=25, offset=0):
query = select(DataPreparation)
if scope_key:
query = query.where(DataPreparation.scope_key == scope_key)
if q:
query = query.where(
or_(
DataPreparation.name.ilike(contains(q), escape="\\"),
DataPreparation.note.ilike(contains(q), escape="\\"),
)
)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(DataPreparation.updated_at.desc(), DataPreparation.id).limit(limit).offset(offset)
)
return page([await self.output(row) for row in rows], total, limit, offset)
async def members(self, preparation_id, q="", dataset_id=None, limit=25, offset=0):
await self.get(preparation_id)
query = select(PreparationField).where(PreparationField.preparation_id == preparation_id)
if dataset_id:
query = query.where(PreparationField.dataset_id == dataset_id)
if q:
query = query.where(
or_(
PreparationField.field_id.ilike(contains(q), escape="\\"),
PreparationField.content["name"].as_string().ilike(contains(q), escape="\\"),
PreparationField.content["description"].as_string().ilike(contains(q), escape="\\"),
)
)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(PreparationField.dataset_id, PreparationField.field_id).limit(limit).offset(offset)
)
return page([row.content for row in rows], total, limit, offset)
async def create(self, name, note, scope, fields):
row = DataPreparation(
id=str(uuid4()), name=name, note=note, scope=scope.model_dump(), scope_key=scope.key()
)
self.db.add(row)
await self.db.flush()
await self.add(row, fields)
return await self.output(row)
async def add(self, row, fields):
for field in fields:
existing = await self.db.get(PreparationField, (row.id, field["id"]))
if existing:
if existing.dataset_id != field["dataset_id"]:
raise HTTPException(422, "同名字段的数据集归属冲突")
continue
self.db.add(
PreparationField(
preparation_id=row.id, field_id=field["id"], dataset_id=field["dataset_id"], content=field
)
)
await self.db.flush()
async def copy_dataset(self, body):
dataset = await Catalog(self.db).dataset(body.scope, body.dataset_id, lock=True)
if not dataset.field_version or dataset.field_version != body.collection_version:
raise HTTPException(409, "数据集尚未完整同步或版本已更新")
batch = await self.db.get(CatalogBatch, dataset.field_version)
if not batch.complete:
raise HTTPException(409, "数据集尚未完整同步")
source = await Catalog(self.db).detail(body.scope, body.dataset_id)
rows = await self.db.scalars(
select(CatalogEntry)
.where(CatalogEntry.batch_id == dataset.field_version)
.order_by(CatalogEntry.id)
)
fields = [self.local_field(row, body.dataset_id, source["name"], body.scope) for row in rows]
return await self.create(
f"{source['name'] or body.dataset_id} · {now():%Y%m%d-%H%M%S-%f}", "", body.scope, fields
)
async def remove(self, refs):
rows = [await self.get(ref.id, ref.version, lock=True) for ref in sorted(refs, key=lambda r: r.id)]
for row in rows:
await self.db.execute(delete(PreparationField).where(PreparationField.preparation_id == row.id))
await self.db.delete(row)
return {"deleted": len(rows)}
async def freeze(self, refs):
"""Lock collection versions and capture source-independent research snapshots atomically."""
snapshots = []
for ref in sorted(refs, key=lambda r: r.id):
row = await self.get(ref.id, ref.version, lock=True)
existing = await self.db.scalar(
select(ResearchInputSnapshot).where(
ResearchInputSnapshot.preparation_id == row.id,
ResearchInputSnapshot.preparation_version == row.version,
)
)
if existing:
snapshots.append(await self.snapshot(existing.id))
continue
fields = [
r.content
for r in await self.db.scalars(
select(PreparationField)
.where(PreparationField.preparation_id == row.id)
.order_by(PreparationField.field_id)
)
]
if not fields:
raise HTTPException(422, "空集合不能用于研究")
fixed = ResearchInputSnapshot(
id=str(uuid4()),
preparation_id=row.id,
preparation_version=row.version,
content=dict(
name=row.name,
scope=row.scope,
fields=fields,
field_ids=[f["id"] for f in fields],
field_types={f["id"]: f["field_type"] for f in fields},
dataset_ids=sorted({f["dataset_id"] for f in fields}),
),
)
self.db.add(fixed)
await self.db.flush()
snapshots.append(await self.snapshot(fixed.id))
return snapshots
async def snapshot(self, snapshot_id):
row = await self.db.get(ResearchInputSnapshot, snapshot_id)
if not row:
raise HTTPException(404, "研究输入快照不存在")
return encode_snapshot(
dict(
**row.content,
id=row.id,
preparation_id=row.preparation_id,
preparation_version=row.preparation_version,
created_at=row.created_at,
)
)
async def bind(self, body):
refs = getattr(body, "preparation_refs", [])
if refs:
fixed = await self.freeze(refs)
body.input_ids = list(dict.fromkeys([*body.input_ids, *[r["id"] for r in fixed]]))
body.preparation_refs = []
if not body.input_ids:
raise HTTPException(422, "请选择非空的数据准备集合")
limit = next(
m.max_length for m in type(body).model_fields["input_ids"].metadata if hasattr(m, "max_length")
)
if len(body.input_ids) > limit:
raise HTTPException(422, f"最多可选择 {limit} 个研究输入")
return body