merge: integrate backtests with main dataset catalog and sequence migration 0004
This commit is contained in:
@@ -35,7 +35,7 @@ class ModelSettingsInput(Contract):
|
||||
|
||||
|
||||
class PageContext(Contract):
|
||||
page: Literal["alphas", "account", "backtests"] = "alphas"
|
||||
page: Literal["alphas", "account", "datasets", "backtests"] = "alphas"
|
||||
backtest_run_id: str | None = Field(default=None, max_length=36)
|
||||
backtest_preview_id: str | None = Field(default=None, max_length=36)
|
||||
backtest_draft_id: str | None = Field(default=None, max_length=36)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Scope-isolated data catalog and immutable template input preparation."""
|
||||
@@ -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]
|
||||
@@ -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)
|
||||
@@ -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]
|
||||
@@ -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
|
||||
@@ -251,6 +251,10 @@ class Runner:
|
||||
await self.ensure_connected(force=kind == "connect")
|
||||
if kind in ("connect", "profile"):
|
||||
await self.refresh_profile()
|
||||
elif kind in ("catalog_sync", "field_sync"):
|
||||
from .catalog.sync import sync_catalog
|
||||
|
||||
await sync_catalog(self, job_id, payload)
|
||||
elif kind == "full_sync":
|
||||
await self.sync_all(job_id)
|
||||
else:
|
||||
|
||||
@@ -18,6 +18,7 @@ from .ai.runtime import AIRuntime
|
||||
from .alphas import list_statement, sorted_statement
|
||||
from .backtests.routes import router as backtest_router
|
||||
from .business import Business, notify_job
|
||||
from .catalog.routes import router as catalog_router
|
||||
from .config import Settings
|
||||
from .db import create_database
|
||||
from .jobs import AUTH_KINDS, Runner, create_job
|
||||
@@ -386,5 +387,6 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
|
||||
app.include_router(backtest_router)
|
||||
app.include_router(api)
|
||||
app.include_router(catalog_router)
|
||||
app.include_router(ai_router(ai_runtime))
|
||||
return app
|
||||
|
||||
@@ -310,3 +310,69 @@ class BacktestEvent(Base):
|
||||
kind: Mapped[str] = mapped_column(String(50))
|
||||
payload: Mapped[dict] = mapped_column(JSON)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class CatalogScope(Base):
|
||||
__tablename__ = "catalog_scopes"
|
||||
key: Mapped[str] = mapped_column(String(200), primary_key=True)
|
||||
scope: Mapped[dict] = mapped_column(JSON)
|
||||
catalog_version: Mapped[str | None] = mapped_column(String(36))
|
||||
synced_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
|
||||
|
||||
class CatalogBatch(Base):
|
||||
__tablename__ = "catalog_batches"
|
||||
id: Mapped[str] = mapped_column(ForeignKey("sync_jobs.id"), primary_key=True)
|
||||
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True)
|
||||
dataset_id: Mapped[str | None] = mapped_column(String(200))
|
||||
complete: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
|
||||
|
||||
class CatalogDataset(Base):
|
||||
__tablename__ = "catalog_datasets"
|
||||
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), primary_key=True)
|
||||
id: Mapped[str] = mapped_column(String(200), primary_key=True)
|
||||
field_version: Mapped[str | None] = mapped_column(ForeignKey("catalog_batches.id"))
|
||||
|
||||
|
||||
class CatalogEntry(Base):
|
||||
"""Immutable published snapshots; staging rows remain invisible until batch completion."""
|
||||
__tablename__ = "catalog_entries"
|
||||
batch_id: Mapped[str] = mapped_column(ForeignKey("catalog_batches.id"), primary_key=True)
|
||||
id: Mapped[str] = mapped_column(String(200), primary_key=True)
|
||||
name: Mapped[str | None] = mapped_column(Text)
|
||||
category: Mapped[str | None] = mapped_column(String(200))
|
||||
subcategory: Mapped[str | None] = mapped_column(String(200))
|
||||
field_type: Mapped[str | None] = mapped_column(String(100))
|
||||
coverage: Mapped[float | None] = mapped_column(Float)
|
||||
user_count: Mapped[int | None] = mapped_column(Integer)
|
||||
alpha_count: Mapped[int | None] = mapped_column(Integer)
|
||||
field_count: Mapped[int | None] = mapped_column(Integer)
|
||||
description: Mapped[str | None] = mapped_column(Text)
|
||||
unit: Mapped[str | None] = mapped_column(Text)
|
||||
synced_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class CatalogNote(Base):
|
||||
__tablename__ = "catalog_notes"
|
||||
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), primary_key=True)
|
||||
dataset_id: Mapped[str] = mapped_column(String(200), primary_key=True)
|
||||
# Empty field_id denotes the dataset; platform identifiers cannot be empty.
|
||||
field_id: Mapped[str] = mapped_column(String(200), primary_key=True, default="")
|
||||
note: Mapped[str] = mapped_column(Text, default="")
|
||||
version: Mapped[int] = mapped_column(Integer, default=1)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
|
||||
class TemplateInput(Base):
|
||||
__tablename__ = "template_inputs"
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True)
|
||||
scope_key: Mapped[str] = mapped_column(ForeignKey("catalog_scopes.key"), index=True)
|
||||
dataset_id: Mapped[str] = mapped_column(String(200))
|
||||
collection_version: Mapped[str] = mapped_column(ForeignKey("catalog_batches.id"))
|
||||
selection: Mapped[str] = mapped_column(String(20))
|
||||
field_ids: Mapped[list] = mapped_column(JSON)
|
||||
field_types: Mapped[dict] = mapped_column(JSON)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
|
||||
|
||||
@@ -223,6 +223,7 @@ class JobOutput(BaseModel):
|
||||
id: str
|
||||
kind: str
|
||||
status: str
|
||||
payload: dict = Field(default_factory=dict)
|
||||
processed: int
|
||||
failed: int
|
||||
total: int | None
|
||||
|
||||
@@ -351,3 +351,12 @@ class WqClient:
|
||||
|
||||
async def pnl(self, alpha_id):
|
||||
return await self.get(f"/alphas/{alpha_id}/recordsets/pnl")
|
||||
|
||||
async def catalog_page(self, scope, dataset_id, offset):
|
||||
"""Read a single scoped page. IDs are query parameters, never upstream paths."""
|
||||
params = {"instrumentType": scope["instrument_type"], "region": scope["region"],
|
||||
"universe": scope["universe"], "delay": scope["delay"],
|
||||
"limit": 50, "offset": offset}
|
||||
if dataset_id is not None:
|
||||
params["dataset.id"] = dataset_id
|
||||
return await self.get("/data-fields" if dataset_id else "/data-sets", params)
|
||||
|
||||
Reference in New Issue
Block a user