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