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 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user