2026-09-12 01:24:02 +08:00
|
|
|
"""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):
|
2026-09-12 02:13:48 +08:00
|
|
|
text_filters = (filters.q, filters.dataset_id, filters.field_type)
|
|
|
|
|
numeric_filters = (
|
|
|
|
|
filters.coverage_min, filters.coverage_max,
|
|
|
|
|
filters.user_count_min, filters.user_count_max,
|
|
|
|
|
filters.alpha_count_min, filters.alpha_count_max,
|
|
|
|
|
)
|
|
|
|
|
if not any(value and value.strip() for value in text_filters) and not any(
|
|
|
|
|
value is not None for value in numeric_filters
|
|
|
|
|
):
|
|
|
|
|
raise HTTPException(422, "在线查询至少需要关键词、数据集、类型或数值筛选条件")
|
2026-09-12 01:24:02 +08:00
|
|
|
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
|