refactor: unify data preparations and research input snapshots
Deploy production / deploy (push) Successful in 53s

This commit is contained in:
yuxuanhui
2026-09-12 01:24:02 +08:00
parent 849f86fef7
commit 394438e753
82 changed files with 4146 additions and 2076 deletions
+1 -27
View File
@@ -3,7 +3,7 @@
from datetime import datetime, timezone
from typing import Annotated, Literal
from pydantic import AfterValidator, BaseModel, Field, model_validator
from pydantic import AfterValidator, BaseModel, Field
from ..schemas import Contract
@@ -54,20 +54,6 @@ class NoteInput(Contract):
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
@@ -107,18 +93,6 @@ class CatalogPage(BaseModel):
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]
-20
View File
@@ -12,8 +12,6 @@ from .contracts import (
CatalogPage,
CollectionOutput,
EntryOutput,
InputOutput,
InputPreparation,
NoteInput,
NoteOutput,
Scope,
@@ -76,24 +74,6 @@ async def sync(request: Request, body: CatalogJobInput):
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:
+5 -65
View File
@@ -4,7 +4,6 @@ The dataset row serializes collection publication and draft creation on PostgreS
No page filters participate in template input selection.
"""
from datetime import timezone
from uuid import uuid4
from fastapi import HTTPException
@@ -18,7 +17,6 @@ from ..models import (
CatalogNote,
CatalogScope,
Job,
TemplateInput,
now,
)
from ..schemas import JobOutput
@@ -164,13 +162,13 @@ class Catalog:
raise HTTPException(409, "研究备注已被修改;当前草稿已保留,请载入最新记录后重新保存")
return dict(note=body.note, version=body.version + 1, updated_at=now())
async def create_job(self, body):
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 = "field_sync" if body.dataset_id else "catalog_sync"
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(
@@ -190,7 +188,7 @@ class Catalog:
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))
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)
@@ -210,64 +208,6 @@ class Catalog:
)
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]
from ..preparations.service import Preparations
return await Preparations(self.db).snapshot(input_id)
+82 -14
View File
@@ -4,6 +4,7 @@ import asyncio
import math
import re
from urllib.parse import parse_qs, urlparse
from uuid import uuid4
from sqlalchemy import select
@@ -64,14 +65,15 @@ def normalize(raw, dataset_id):
)
async def sync_catalog(runner, job_id, payload):
async def sync_catalog(runner, job_id, payload, *, batch_id=None, full=False):
scope = Scope.model_validate(payload["scope"])
dataset_id = payload.get("dataset_id")
batch_id = batch_id or job_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)
batch = await db.get(CatalogBatch, batch_id)
if batch.complete:
return
offset = batch.offset
while True:
await runner.checkpoint(job_id, {"next_retry_at": None})
raw = await runner.client.catalog_page(scope.model_dump(), dataset_id, offset)
@@ -97,12 +99,12 @@ async def sync_catalog(runner, job_id, payload):
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
batch = await db.get(CatalogBatch, job_id)
batch = await db.get(CatalogBatch, batch_id)
added = 0
for entry in entries:
if await db.get(CatalogEntry, (job_id, entry["id"])):
if await db.get(CatalogEntry, (batch_id, entry["id"])):
continue
db.add(CatalogEntry(batch_id=job_id, **entry))
db.add(CatalogEntry(batch_id=batch_id, **entry))
await db.flush()
added += 1
owner = dataset_id or entry["id"]
@@ -112,25 +114,28 @@ async def sync_catalog(runner, job_id, payload):
if rows and not added:
raise WqError("平台分页重复且未前进,已保留进度", "invalid_response")
batch.count += added
job.processed = batch.count
if not full:
job.processed = batch.count
offset += len(rows)
job.checkpoint = dict(offset=offset, done=not more)
batch.offset = offset
job.checkpoint = {**job.checkpoint, "offset": offset, "done": not more, "current_field_count": batch.count}
job.updated_at = now()
if not more:
batch.complete, batch.completed_at = True, now()
job.total = batch.count
if not full:
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
dataset.field_version = batch_id
else:
scope_row = await db.get(CatalogScope, scope.key())
scope_row.catalog_version, scope_row.synced_at = job_id, now()
scope_row.catalog_version, scope_row.synced_at = batch_id, now()
ids = (
await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id))
await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == batch_id))
).all()
for item_id in ids:
if not await db.get(CatalogDataset, (scope.key(), item_id)):
@@ -138,3 +143,66 @@ async def sync_catalog(runner, job_id, payload):
await db.commit()
if not more:
return
async def sync_full_catalog(runner, job_id, payload):
"""Resume each dataset batch independently; only publish complete enumerations."""
scope = Scope.model_validate(payload["scope"])
options = await runner.client.get_platform_setting_options()
if not any(r["instrument_type"] == scope.instrument_type and r["region"] == scope.region
and r["delay"] == scope.delay and scope.universe in r["universes"]
for r in options["instrument_options"]):
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.checkpoint = {**job.checkpoint, "error_code": "invalid_scope"}
raise WqError("平台不支持该研究范围", "invalid_scope")
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.checkpoint = {**{k: v for k, v in job.checkpoint.items() if k != "error_code"}, "phase": "catalog"}
await sync_catalog(runner, job_id, {"scope": payload["scope"]}, full=True)
async with runner.sessions() as db:
ids = list(await db.scalars(select(CatalogEntry.id).where(CatalogEntry.batch_id == job_id)
.order_by(CatalogEntry.id)))
job = await db.get(Job, job_id)
failures = dict(job.checkpoint.get("failures", {}))
completed = 0
await runner.checkpoint(job_id, {"total": len(ids)})
for dataset_id in ids:
async with runner.sessions.begin() as db:
batch = await db.scalar(select(CatalogBatch).where(CatalogBatch.job_id == job_id,
CatalogBatch.dataset_id == dataset_id))
if not batch:
batch = CatalogBatch(id=str(uuid4()), job_id=job_id, scope_key=scope.key(), dataset_id=dataset_id)
db.add(batch)
await db.flush()
batch_id, complete = batch.id, batch.complete
job = await db.get(Job, job_id)
if job.cancel_requested:
raise asyncio.CancelledError()
job.checkpoint = {**job.checkpoint, "phase": "fields", "dataset_id": dataset_id,
"datasets_completed": completed, "datasets_total": len(ids),
"offset": batch.offset, "current_field_count": batch.count}
if not complete:
try:
await sync_catalog(runner, job_id, {"scope": payload["scope"], "dataset_id": dataset_id},
batch_id=batch_id, full=True)
except WqError as exc:
if exc.code in ("disconnected", "authentication_failed", "identity_mismatch", "verification_required", "network_error"):
raise
failures[dataset_id] = str(exc)
if complete or (await _batch_complete(runner, batch_id)):
failures.pop(dataset_id, None)
completed += 1
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.processed, job.failed = completed, len(failures)
job.checkpoint = {**job.checkpoint, "datasets_completed": completed, "failures": failures}
async with runner.sessions.begin() as db:
job = await db.get(Job, job_id)
job.error = f"{len(failures)} 个数据集同步失败" if failures else None
job.checkpoint = {**job.checkpoint, "phase": "finished"}
async def _batch_complete(runner, batch_id):
async with runner.sessions() as db:
return (await db.get(CatalogBatch, batch_id)).complete