214 lines
9.2 KiB
Python
214 lines
9.2 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 uuid import uuid4
|
|
|
|
from fastapi import HTTPException
|
|
from sqlalchemy import func, or_, select, update
|
|
|
|
from ..models import (
|
|
Account,
|
|
CatalogBatch,
|
|
CatalogDataset,
|
|
CatalogEntry,
|
|
CatalogNote,
|
|
CatalogScope,
|
|
Job,
|
|
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, *, full=False):
|
|
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 = "catalog_full_sync" if full else "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, job_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 input(self, input_id):
|
|
from ..preparations.service import Preparations
|
|
return await Preparations(self.db).snapshot(input_id)
|