feat: implement scoped dataset catalog and template input drafts
This commit is contained in:
@@ -0,0 +1,273 @@
|
||||
"""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]
|
||||
Reference in New Issue
Block a user