Files
worldquant-alpha-system/backend/app/catalog/service.py
T

274 lines
12 KiB
Python

"""Catalog business operations. Callers own authorization and transaction commits.
The dataset row serializes collection publication and draft creation on PostgreSQL.
No page filters participate in template input selection.
"""
from datetime import timezone
from uuid import uuid4
from fastapi import HTTPException
from sqlalchemy import func, or_, select, update
from ..models import (
Account,
CatalogBatch,
CatalogDataset,
CatalogEntry,
CatalogNote,
CatalogScope,
Job,
TemplateInput,
now,
)
from ..schemas import JobOutput
from .contracts import EntryOutput, Scope
class Catalog:
def __init__(self, db):
self.db = db
async def dataset(self, scope, dataset_id, lock=False):
query = select(CatalogDataset).where(
CatalogDataset.scope_key == scope.key(), CatalogDataset.id == dataset_id
)
row = await self.db.scalar(query.with_for_update() if lock else query)
if not row:
raise HTTPException(404, "该范围的数据集尚未同步")
return row
async def search(self, filters, dataset_id=None):
scope = await self.db.get(CatalogScope, filters.key())
version = scope.catalog_version if scope else None
if dataset_id:
version = (await self.dataset(filters, dataset_id)).field_version
batch = await self.db.get(CatalogBatch, version) if version else None
base = (
select(CatalogEntry).where(CatalogEntry.batch_id == version)
if version
else select(CatalogEntry).where(False)
)
query = base
if filters.q:
pattern = "%" + filters.q.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
query = query.where(
or_(
CatalogEntry.id.ilike(pattern, escape="\\"), CatalogEntry.name.ilike(pattern, escape="\\")
)
)
for key in ("category", "subcategory", "field_type"):
value = getattr(filters, key)
if value is not None:
query = query.where(getattr(CatalogEntry, key) == value)
if filters.coverage_min is not None:
query = query.where(CatalogEntry.coverage >= filters.coverage_min)
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
column = getattr(CatalogEntry, filters.sort)
if dataset_id is None and filters.sort == "field_count":
published_count = (
select(CatalogBatch.count)
.join(CatalogDataset, CatalogDataset.field_version == CatalogBatch.id)
.where(CatalogDataset.scope_key == filters.key(), CatalogDataset.id == CatalogEntry.id)
.correlate(CatalogEntry)
.scalar_subquery()
)
column = func.coalesce(published_count, CatalogEntry.field_count)
query = query.order_by(
(column.desc() if filters.direction == "desc" else column.asc()).nulls_last(), CatalogEntry.id
)
entries = (await self.db.scalars(query.limit(filters.limit).offset(filters.offset))).all()
items = [EntryOutput.model_validate(e, from_attributes=True).model_dump() for e in entries]
if not dataset_id and items:
datasets = (
await self.db.scalars(
select(CatalogDataset).where(
CatalogDataset.scope_key == filters.key(),
CatalogDataset.id.in_([i["id"] for i in items]),
)
)
).all()
versions = {d.id: d.field_version for d in datasets}
batches = (
await self.db.scalars(
select(CatalogBatch).where(CatalogBatch.id.in_([v for v in versions.values() if v]))
)
).all()
counts = {b.id: b.count for b in batches}
for item in items:
item["collection_version"] = versions.get(item["id"])
item["complete_count"] = counts.get(versions.get(item["id"]))
categories = {}
for category, subcategory in (
await self.db.execute(
base.with_only_columns(CatalogEntry.category, CatalogEntry.subcategory).distinct()
)
).all():
if category:
categories.setdefault(category, [])
if subcategory and subcategory not in categories[category]:
categories[category].append(subcategory)
types = (
await self.db.scalars(
base.with_only_columns(CatalogEntry.field_type)
.where(CatalogEntry.field_type.is_not(None))
.distinct()
.order_by(CatalogEntry.field_type)
)
).all()
return dict(
items=items,
total=total,
limit=filters.limit,
offset=filters.offset,
collection_version=version,
complete_count=batch.count if batch else None,
synced_at=batch.completed_at if batch else None,
categories=categories,
field_types=types,
)
async def detail(self, scope, dataset_id, field_id=""):
dataset = await self.dataset(scope, dataset_id)
scope_row = await self.db.get(CatalogScope, scope.key())
version = dataset.field_version if field_id else scope_row.catalog_version
entry = await self.db.get(CatalogEntry, (version, field_id or dataset_id)) if version else None
if not entry:
raise HTTPException(404, "该范围的对象尚未完整同步")
note = await self.db.get(CatalogNote, (scope.key(), dataset_id, field_id))
batch = await self.db.get(CatalogBatch, dataset.field_version) if dataset.field_version else None
return dict(
**EntryOutput.model_validate(entry, from_attributes=True).model_dump(
exclude={"research", "scope", "dataset_id", "collection_version", "complete_count"}
),
research=dict(note=note.note, version=note.version, updated_at=note.updated_at),
scope=Scope.model_validate(scope.model_dump(include=set(Scope.model_fields))),
dataset_id=dataset_id,
collection_version=dataset.field_version,
complete_count=batch.count if batch else None,
)
async def save_note(self, scope, dataset_id, field_id, body):
await self.detail(scope, dataset_id, field_id)
result = await self.db.execute(
update(CatalogNote)
.where(
CatalogNote.scope_key == scope.key(),
CatalogNote.dataset_id == dataset_id,
CatalogNote.field_id == field_id,
CatalogNote.version == body.version,
)
.values(note=body.note, version=CatalogNote.version + 1, updated_at=now())
)
if result.rowcount != 1:
raise HTTPException(409, "研究备注已被修改;当前草稿已保留,请载入最新记录后重新保存")
return dict(note=body.note, version=body.version + 1, updated_at=now())
async def create_job(self, body):
account = await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
if not account.password_encrypted or account.connection_status in ("disconnected", "error"):
raise HTTPException(409, "请先连接 WorldQuant")
if body.dataset_id:
await self.dataset(body.scope, body.dataset_id)
kind = "field_sync" if body.dataset_id else "catalog_sync"
payload = body.model_dump(mode="json")
jobs = (
await self.db.scalars(
select(Job).where(
Job.kind == kind,
Job.status.in_(("queued", "running", "waiting_auth", "waiting_connection")),
)
)
).all()
for job in jobs:
if job.payload == payload:
return JobOutput.model_validate(job)
scope = await self.db.get(CatalogScope, body.scope.key())
if not scope:
self.db.add(CatalogScope(key=body.scope.key(), scope=body.scope.model_dump()))
await self.db.flush()
job = Job(id=str(uuid4()), kind=kind, payload=payload)
self.db.add(job)
await self.db.flush()
self.db.add(CatalogBatch(id=job.id, scope_key=body.scope.key(), dataset_id=body.dataset_id))
await self.db.flush()
return JobOutput.model_validate(job)
async def collection(self, scope, dataset_id):
"""Return membership only for the published collection, independent of table filters."""
dataset = await self.dataset(scope, dataset_id)
ids = []
if dataset.field_version:
ids = list(
(
await self.db.scalars(
select(CatalogEntry.id)
.where(CatalogEntry.batch_id == dataset.field_version)
.order_by(CatalogEntry.id)
)
).all()
)
return dict(collection_version=dataset.field_version, field_ids=ids)
async def prepare(self, body):
dataset = await self.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 or batch.scope_key != body.scope.key() or batch.dataset_id != body.dataset_id:
raise HTTPException(409, "字段集合不完整")
entries = (
await self.db.scalars(
select(CatalogEntry).where(CatalogEntry.batch_id == batch.id).order_by(CatalogEntry.id)
)
).all()
fields = {e.id: e.field_type for e in entries}
excluded = set(body.excluded_ids)
if excluded - fields.keys():
raise HTTPException(422, "排除项含未知、跨范围或其他数据集字段")
chosen = {key: value for key, value in fields.items() if key not in excluded}
if not chosen:
raise HTTPException(422, "模板输入至少需要一个字段")
row = TemplateInput(
id=str(uuid4()),
scope_key=body.scope.key(),
dataset_id=body.dataset_id,
collection_version=batch.id,
selection=body.selection,
field_ids=list(chosen),
field_types=chosen,
)
self.db.add(row)
await self.db.flush()
return await self.input(row.id)
async def input(self, input_id):
row = await self.db.get(TemplateInput, input_id)
if not row:
raise HTTPException(404, "输入草稿不存在")
scope = await self.db.get(CatalogScope, row.scope_key)
return dict(
id=row.id,
status="draft",
scope=scope.scope,
dataset_id=row.dataset_id,
collection_version=row.collection_version,
selection=row.selection,
field_ids=row.field_ids,
field_types=row.field_types,
created_at=row.created_at.replace(tzinfo=timezone.utc)
if row.created_at.tzinfo is None
else row.created_at,
)
async def inputs(self, scope):
ids = (
await self.db.scalars(
select(TemplateInput.id)
.where(TemplateInput.scope_key == scope.key())
.order_by(TemplateInput.created_at.desc())
.limit(100)
)
).all()
return [await self.input(i) for i in ids]