"""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