feat: implement scoped dataset catalog and template input drafts

This commit is contained in:
yuxuanhui
2026-09-08 09:22:53 +08:00
parent 404a4d8a04
commit 01d169d118
30 changed files with 3024 additions and 37 deletions
+1
View File
@@ -0,0 +1 @@
"""Scope-isolated data catalog and immutable template input preparation."""
+137
View File
@@ -0,0 +1,137 @@
"""Explicit research scope and catalog contracts; unknown platform types remain strings."""
from datetime import datetime, timezone
from typing import Annotated, Literal
from pydantic import AfterValidator, BaseModel, Field, model_validator
from ..schemas import Contract
def utc_timestamp(value: datetime) -> datetime:
"""SQLite drops tzinfo; catalog source times always denote UTC instants."""
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value
UTCTimestamp = Annotated[datetime, AfterValidator(utc_timestamp)]
# Supported research scopes, not an assertion about a connected account's permissions.
UNIVERSES = {
"USA": ["TOP3000", "TOP1000", "TOP500", "TOP200"],
"CHN": ["TOP2000"],
"EUR": ["TOP2500", "TOP1200"],
"ASI": ["TOP1000"],
"GLB": ["TOP3000"],
"JPN": ["TOP1600"],
"HKG": ["TOP800"],
}
class Scope(Contract):
instrument_type: Literal["EQUITY"] = "EQUITY"
region: str
universe: str
delay: int = Field(ge=0, le=1)
@model_validator(mode="after")
def valid_scope(self):
if self.universe not in UNIVERSES.get(self.region, []):
raise ValueError("不支持的 Region / Universe 组合")
return self
def key(self):
return f"{self.instrument_type}|{self.region}|{self.universe}|{self.delay}"
class CatalogFilters(Scope):
q: str = Field(default="", max_length=300)
category: str | None = None
subcategory: str | None = None
field_type: str | None = None
coverage_min: float | None = Field(default=None, ge=0, le=1)
sort: Literal[
"id", "name", "category", "field_count", "coverage", "user_count", "alpha_count", "field_type"
] = "name"
direction: Literal["asc", "desc"] = "asc"
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class CatalogJobInput(Contract):
scope: Scope
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
class NoteInput(Contract):
note: str = Field(max_length=20000)
version: int = Field(ge=1)
class InputPreparation(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
collection_version: str
selection: Literal["all", "explicit"] = "all"
excluded_ids: list[str] = Field(default_factory=list, max_length=100000)
@model_validator(mode="after")
def valid_selection(self):
if self.selection == "all" and self.excluded_ids:
raise ValueError("全部字段不能同时提供排除项")
return self
class NoteOutput(BaseModel):
note: str
version: int
updated_at: UTCTimestamp
class EntryOutput(BaseModel):
id: str
name: str | None
category: str | None
subcategory: str | None
field_type: str | None
coverage: float | None
user_count: int | None
alpha_count: int | None
field_count: int | None
description: str | None
unit: str | None
synced_at: UTCTimestamp
collection_version: str | None = None
complete_count: int | None = None
research: NoteOutput | None = None
scope: Scope | None = None
dataset_id: str | None = None
class CatalogPage(BaseModel):
items: list[EntryOutput]
total: int
limit: int
offset: int
collection_version: str | None
complete_count: int | None
synced_at: UTCTimestamp | None
categories: dict[str, list[str]] = Field(default_factory=dict)
field_types: list[str] = Field(default_factory=list)
class InputOutput(BaseModel):
id: str
status: Literal["draft"] = "draft"
scope: Scope
dataset_id: str
collection_version: str
selection: str
field_ids: list[str]
field_types: dict[str, str | None]
created_at: UTCTimestamp
class CollectionOutput(BaseModel):
collection_version: str | None
field_ids: list[str]
+99
View File
@@ -0,0 +1,99 @@
"""Authenticated catalog endpoints; writes inherit the application origin guard."""
from typing import Annotated
from fastapi import APIRouter, Depends, Query, Request
from ..schemas import JobOutput
from ..security import require_auth
from .contracts import (
UNIVERSES,
CatalogFilters,
CatalogJobInput,
CatalogPage,
CollectionOutput,
EntryOutput,
InputOutput,
InputPreparation,
NoteInput,
NoteOutput,
Scope,
)
from .service import Catalog
router = APIRouter(prefix="/api/v1/catalog", tags=["catalog"], dependencies=[Depends(require_auth)])
@router.get("/scopes")
async def scopes() -> dict[str, list[str]]:
return UNIVERSES
@router.get("/datasets", response_model=CatalogPage)
async def datasets(request: Request, filters: Annotated[CatalogFilters, Query()]):
async with request.app.state.sessions() as db:
return await Catalog(db).search(filters)
@router.get("/datasets/{dataset_id}", response_model=EntryOutput)
async def detail(request: Request, dataset_id: str, scope: Annotated[Scope, Query()]):
async with request.app.state.sessions() as db:
return await Catalog(db).detail(scope, dataset_id)
@router.get("/datasets/{dataset_id}/fields", response_model=CatalogPage)
async def fields(request: Request, dataset_id: str, filters: Annotated[CatalogFilters, Query()]):
async with request.app.state.sessions() as db:
return await Catalog(db).search(filters, dataset_id)
@router.get("/datasets/{dataset_id}/fields/{field_id}", response_model=EntryOutput)
async def field(request: Request, dataset_id: str, field_id: str, scope: Annotated[Scope, Query()]):
async with request.app.state.sessions() as db:
return await Catalog(db).detail(scope, dataset_id, field_id)
@router.patch("/datasets/{dataset_id}/research", response_model=NoteOutput)
async def note(request: Request, dataset_id: str, scope: Annotated[Scope, Query()], body: NoteInput):
async with request.app.state.sessions.begin() as db:
return await Catalog(db).save_note(scope, dataset_id, "", body)
@router.patch("/datasets/{dataset_id}/fields/{field_id}/research", response_model=NoteOutput)
async def field_note(
request: Request, dataset_id: str, field_id: str, scope: Annotated[Scope, Query()], body: NoteInput
):
async with request.app.state.sessions.begin() as db:
return await Catalog(db).save_note(scope, dataset_id, field_id, body)
@router.post("/sync-jobs", status_code=202, response_model=JobOutput)
async def sync(request: Request, body: CatalogJobInput):
async with request.app.state.sessions.begin() as db:
result = await Catalog(db).create_job(body)
request.app.state.runner.wake.set()
return result
@router.post("/inputs", status_code=201, response_model=InputOutput)
async def prepare(request: Request, body: InputPreparation):
async with request.app.state.sessions.begin() as db:
return await Catalog(db).prepare(body)
@router.get("/inputs", response_model=list[InputOutput])
async def inputs(request: Request, scope: Annotated[Scope, Query()]):
async with request.app.state.sessions() as db:
return await Catalog(db).inputs(scope)
@router.get("/inputs/{input_id}", response_model=InputOutput)
async def get_input(request: Request, input_id: str):
async with request.app.state.sessions() as db:
return await Catalog(db).input(input_id)
@router.get("/datasets/{dataset_id}/collection", response_model=CollectionOutput)
async def collection(request: Request, dataset_id: str, scope: Annotated[Scope, Query()]):
async with request.app.state.sessions() as db:
return await Catalog(db).collection(scope, dataset_id)
+273
View File
@@ -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]
+139
View File
@@ -0,0 +1,139 @@
"""Publish complete enumerations only; retain staging checkpoints and old versions."""
import asyncio
import math
import re
from urllib.parse import parse_qs, urlparse
from sqlalchemy import select
from ..models import CatalogBatch, CatalogDataset, CatalogEntry, CatalogNote, CatalogScope, Job, now
from ..worldquant import WqError
from .contracts import Scope
def identifier(value):
if not isinstance(value, str) or not re.fullmatch(r"[A-Za-z0-9_.-]{1,200}", value):
raise WqError("平台目录包含无法识别的 ID,已保留进度", "invalid_response")
return value
def label(value):
if isinstance(value, dict):
value = value.get("name") or value.get("id")
return value if isinstance(value, str) and value else None
def number(value, integer=False):
if (
isinstance(value, bool)
or not isinstance(value, (int, float))
or not math.isfinite(value)
or value < 0
):
return None
return int(value) if integer and value == int(value) else None if integer else value
def normalize(raw, dataset_id):
if not isinstance(raw, dict):
raise WqError("平台目录记录格式无法识别", "invalid_response")
item_id = identifier(raw.get("id"))
owner = raw.get("dataset")
owner = owner.get("id") if isinstance(owner, dict) else owner
if dataset_id and owner != dataset_id:
raise WqError("平台返回了其他数据集的字段", "invalid_response")
coverage = number(raw.get("coverage"))
# BRAIN coverage is a fraction. Never guess that a value >1 means percent.
# Real-account schema/units still require read-only integration verification.
if coverage is not None and coverage > 1:
raise WqError("平台覆盖率单位无法确认,应为 0–1", "invalid_response")
return dict(
id=item_id,
name=label(raw.get("name")) or item_id,
category=label(raw.get("category")),
subcategory=label(raw.get("subcategory")),
field_type=label(raw.get("type")) if dataset_id else None,
coverage=coverage,
user_count=number(raw.get("userCount"), True),
alpha_count=number(raw.get("alphaCount"), True),
field_count=number(raw.get("fieldCount"), True),
description=label(raw.get("description")),
unit=label(raw.get("unit")),
)
async def sync_catalog(runner, job_id, payload):
scope = Scope.model_validate(payload["scope"])
dataset_id = payload.get("dataset_id")
async with runner.sessions() as db:
checkpoint = (await db.get(Job, job_id)).checkpoint
if checkpoint.get("done"):
return
offset = checkpoint.get("offset", 0)
while True:
await runner.checkpoint(job_id, {"next_retry_at": None})
raw = await runner.client.catalog_page(scope.model_dump(), dataset_id, offset)
rows = raw.get("results")
if not isinstance(rows, list):
raise WqError("平台目录缺少 results,已保留进度", "invalid_response")
entries = [normalize(r, dataset_id) for r in rows]
# Always probe to exhaustion if next is absent; count alone cannot prove completeness.
next_page = raw.get("next")
if "next" in raw and next_page is not None:
if not isinstance(next_page, str) or not next_page:
raise WqError("平台 next 分页格式无法识别", "invalid_response")
parsed = urlparse(next_page)
expected_path = "/data-fields" if dataset_id else "/data-sets"
offsets = parse_qs(parsed.query).get("offset", [])
if parsed.path.rstrip("/") != expected_path or offsets != [str(offset + len(rows))]:
raise WqError("平台 next 分页未按预期前进", "invalid_response")
more = next_page is not None if "next" in raw else bool(rows)
count = number(raw.get("count"), True)
if (more and not rows) or (not more and count is not None and offset + len(rows) < count):
raise WqError("平台分页提前结束,未发布不完整集合", "invalid_response")
async with runner.sessions() as db:
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
batch = await db.get(CatalogBatch, job_id)
added = 0
for entry in entries:
if await db.get(CatalogEntry, (job_id, entry["id"])):
continue
db.add(CatalogEntry(batch_id=job_id, **entry))
await db.flush()
added += 1
owner = dataset_id or entry["id"]
field_id = entry["id"] if dataset_id else ""
if not await db.get(CatalogNote, (scope.key(), owner, field_id)):
db.add(CatalogNote(scope_key=scope.key(), dataset_id=owner, field_id=field_id))
if rows and not added:
raise WqError("平台分页重复且未前进,已保留进度", "invalid_response")
batch.count += added
job.processed = batch.count
offset += len(rows)
job.checkpoint = dict(offset=offset, done=not more)
job.updated_at = now()
if not more:
batch.complete, batch.completed_at = True, now()
job.total = batch.count
if dataset_id:
dataset = await db.scalar(
select(CatalogDataset)
.where(CatalogDataset.scope_key == scope.key(), CatalogDataset.id == dataset_id)
.with_for_update()
)
dataset.field_version = job_id
else:
scope_row = await db.get(CatalogScope, scope.key())
scope_row.catalog_version, scope_row.synced_at = job_id, now()
ids = (
await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id))
).all()
for item_id in ids:
if not await db.get(CatalogDataset, (scope.key(), item_id)):
db.add(CatalogDataset(scope_key=scope.key(), id=item_id))
await db.commit()
if not more:
return