Files
worldquant-alpha-system/backend/app/preparations/routes.py
T

196 lines
7.6 KiB
Python
Raw Normal View History

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