198 lines
7.7 KiB
Python
198 lines
7.7 KiB
Python
"""Authenticated preparation and field-directory endpoints."""
|
|
|
|
from typing import Annotated, Literal
|
|
|
|
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),
|
|
sort: Literal["name", "created_at", "updated_at"] = "updated_at",
|
|
direction: Literal["asc", "desc"] = "desc",
|
|
):
|
|
async with request.app.state.sessions() as db:
|
|
return await Preparations(db).list(q, scope_key, limit, offset, sort, direction)
|
|
|
|
|
|
@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},
|
|
}
|