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 -1
View File
@@ -35,7 +35,7 @@ class ModelSettingsInput(Contract):
class PageContext(Contract):
page: Literal["alphas", "account"] = "alphas"
page: Literal["alphas", "account", "datasets"] = "alphas"
alpha_id: str | None = Field(default=None, max_length=100, pattern=r"^[A-Za-z0-9_-]+$")
selected_ids: list[str] = Field(default_factory=list, max_length=100)
filters: AlphaFilters = Field(default_factory=AlphaFilters)
+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
+4
View File
@@ -245,6 +245,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:
+2
View File
@@ -17,6 +17,7 @@ from .ai.routes import router as ai_router
from .ai.runtime import AIRuntime
from .alphas import list_statement, sorted_statement
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
@@ -381,5 +382,6 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
return result
app.include_router(api)
app.include_router(catalog_router)
app.include_router(ai_router(ai_runtime))
return app
+66
View File
@@ -199,3 +199,69 @@ class AIToolCall(Base):
status: Mapped[str] = mapped_column(String(30), default="pending")
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=now)
__table_args__ = (UniqueConstraint("run_id", "call_id"),)
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)
+1
View File
@@ -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
+9
View File
@@ -277,3 +277,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)
@@ -0,0 +1,92 @@
"""scope catalog collections notes and input drafts"""
from alembic import op
import sqlalchemy as sa
revision = '0003'
down_revision = '0002'
branch_labels = None
depends_on = None
def upgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.create_table('catalog_scopes',
sa.Column('key', sa.String(length=200), nullable=False),
sa.Column('scope', sa.JSON(), nullable=False),
sa.Column('catalog_version', sa.String(length=36), nullable=True),
sa.Column('synced_at', sa.DateTime(timezone=True), nullable=True),
sa.PrimaryKeyConstraint('key')
)
op.create_table('catalog_batches',
sa.Column('id', sa.String(length=36), nullable=False),
sa.Column('scope_key', sa.String(length=200), nullable=False),
sa.Column('dataset_id', sa.String(length=200), nullable=True),
sa.Column('complete', sa.Boolean(), nullable=False),
sa.Column('count', sa.Integer(), nullable=False),
sa.Column('completed_at', sa.DateTime(timezone=True), nullable=True),
sa.ForeignKeyConstraint(['id'], ['sync_jobs.id'], ),
sa.ForeignKeyConstraint(['scope_key'], ['catalog_scopes.key'], ),
sa.PrimaryKeyConstraint('id')
)
op.create_index(op.f('ix_catalog_batches_scope_key'), 'catalog_batches', ['scope_key'], unique=False)
op.create_table('catalog_notes',
sa.Column('scope_key', sa.String(length=200), nullable=False),
sa.Column('dataset_id', sa.String(length=200), nullable=False),
sa.Column('field_id', sa.String(length=200), nullable=False),
sa.Column('note', sa.Text(), nullable=False),
sa.Column('version', sa.Integer(), nullable=False),
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(['scope_key'], ['catalog_scopes.key'], ),
sa.PrimaryKeyConstraint('scope_key', 'dataset_id', 'field_id')
)
op.create_table('catalog_datasets',
sa.Column('scope_key', sa.String(length=200), nullable=False),
sa.Column('id', sa.String(length=200), nullable=False),
sa.Column('field_version', sa.String(length=36), nullable=True),
sa.ForeignKeyConstraint(['field_version'], ['catalog_batches.id'], ),
sa.ForeignKeyConstraint(['scope_key'], ['catalog_scopes.key'], ),
sa.PrimaryKeyConstraint('scope_key', 'id')
)
op.create_table('catalog_entries',
sa.Column('batch_id', sa.String(length=36), nullable=False),
sa.Column('id', sa.String(length=200), nullable=False),
sa.Column('name', sa.Text(), nullable=True),
sa.Column('category', sa.String(length=200), nullable=True),
sa.Column('subcategory', sa.String(length=200), nullable=True),
sa.Column('field_type', sa.String(length=100), nullable=True),
sa.Column('coverage', sa.Float(), nullable=True),
sa.Column('user_count', sa.Integer(), nullable=True),
sa.Column('alpha_count', sa.Integer(), nullable=True),
sa.Column('field_count', sa.Integer(), nullable=True),
sa.Column('description', sa.Text(), nullable=True),
sa.Column('unit', sa.Text(), nullable=True),
sa.Column('synced_at', sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(['batch_id'], ['catalog_batches.id'], ),
sa.PrimaryKeyConstraint('batch_id', 'id')
)
op.create_table('template_inputs',
sa.Column('id', sa.String(length=36), nullable=False),
sa.Column('scope_key', sa.String(length=200), nullable=False),
sa.Column('dataset_id', sa.String(length=200), nullable=False),
sa.Column('collection_version', sa.String(length=36), nullable=False),
sa.Column('selection', sa.String(length=20), nullable=False),
sa.Column('field_ids', sa.JSON(), nullable=False),
sa.Column('field_types', sa.JSON(), nullable=False),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(['collection_version'], ['catalog_batches.id'], ),
sa.ForeignKeyConstraint(['scope_key'], ['catalog_scopes.key'], ),
sa.PrimaryKeyConstraint('id')
)
op.create_index(op.f('ix_template_inputs_scope_key'), 'template_inputs', ['scope_key'], unique=False)
# ### end Alembic commands ###
def downgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.drop_index(op.f('ix_template_inputs_scope_key'), table_name='template_inputs')
op.drop_table('template_inputs')
op.drop_table('catalog_entries')
op.drop_table('catalog_datasets')
op.drop_table('catalog_notes')
op.drop_index(op.f('ix_catalog_batches_scope_key'), table_name='catalog_batches')
op.drop_table('catalog_batches')
op.drop_table('catalog_scopes')
# ### end Alembic commands ###
+4
View File
@@ -12,6 +12,7 @@ from app.main import create_app
from app.models import Base
from app.worldquant import WqClient
from tests.ai_fake import fake_model
from tests.catalog_fake import catalog_response
TEST_PASSWORD = "browser-test-password"
@@ -104,6 +105,9 @@ def create_test_app():
)
if request.method != "GET":
raise AssertionError("Browser acceptance attempted an upstream mutation")
catalog = catalog_response(request)
if catalog is not None:
return catalog
if path == "/users/self":
return httpx.Response(
200,
+55
View File
@@ -0,0 +1,55 @@
"""Synthetic HTTP catalog, including page overlap and unknown metrics."""
import httpx
def field_records(dataset="TEST_FIN", count=123):
return [
dict(
id=f"{dataset}_{i:03}",
name=f"TEST 字段 {i:03}",
dataset={"id": dataset},
type="FUTURE_TYPE" if i == 122 else "VECTOR" if i % 3 == 0 else "MATRIX",
coverage=None if i == 122 else 0.95 if i % 2 else 0.6,
userCount=None if i == 122 else i,
alphaCount=i * 2,
description=None if i == 122 else f"合成字段说明 {i}",
)
for i in range(count)
]
def catalog_response(request, fields=None):
path, params = request.url.path, request.url.params
if path not in ("/data-sets", "/data-fields"):
return None
assert request.method == "GET"
assert params["instrumentType"] == "EQUITY"
assert params["region"] and params["universe"] and params["delay"] in ("0", "1")
dataset = params.get("dataset.id", "TEST_FIN")
rows = (
[
{
"id": "TEST_FIN",
"name": "TEST 财务报表",
"category": {"name": "基本面"},
"subcategory": {"name": "财务报表"},
"fieldCount": 123,
"description": "合成数据,仅用于验收",
},
{
"id": "TEST_NEWS",
"name": "TEST 新闻",
"category": {"name": "新闻"},
"subcategory": {"name": "情绪"},
"fieldCount": 3,
},
{"id": "TEST_UNKNOWN", "name": "TEST 未分类", "fieldCount": 0},
]
if path == "/data-sets"
else (fields if fields is not None else field_records(dataset, 123 if dataset == "TEST_FIN" else 3))
)
if path == "/data-fields" and len(rows) > 50:
rows = rows[:50] + [rows[49]] + rows[50:]
offset, limit = int(params.get("offset", 0)), int(params.get("limit", 50))
return httpx.Response(200, json={"results": rows[offset : offset + limit]})
+124
View File
@@ -0,0 +1,124 @@
"""One-off acceptance against the dedicated local PostgreSQL catalog_test database."""
import asyncio
import os
import re
from alembic import command
from alembic.config import Config
from cryptography.fernet import Fernet
from sqlalchemy import text
from sqlalchemy.ext.asyncio import create_async_engine
database_name = os.environ.get("WQ_CATALOG_ACCEPTANCE_DATABASE", "catalog_flow_test")
if not re.fullmatch(r"catalog_[a-z0-9_]{1,40}", database_name):
raise ValueError("Acceptance requires a dedicated catalog_* database")
URL = f"postgresql+asyncpg://postgres:catalog-test-only@127.0.0.1:18436/{database_name}"
os.environ.update(
DATABASE_URL=URL, ADMIN_PASSWORD="migration-test-only", ENCRYPTION_KEY=Fernet.generate_key().decode()
)
async def sql(statement):
engine = create_async_engine(URL)
async with engine.begin() as connection:
result = await connection.execute(text(statement))
value = result.fetchall() if result.returns_rows else None
await engine.dispose()
return value
if __name__ == "__main__":
config = Config("alembic.ini")
if asyncio.run(sql("SELECT tablename FROM pg_tables WHERE schemaname='public'")):
raise RuntimeError("Acceptance database must be empty; existing data will not be overwritten")
command.upgrade(config, "0002")
asyncio.run(
sql(
"INSERT INTO alphas (id, hidden, settings, is_metrics, os_metrics, checks, synced_at, raw) VALUES ('MIGRATION_TEST', false, '{}', '{}', '{}', '[]', now(), '{}');"
)
)
asyncio.run(
sql(
"INSERT INTO research (alpha_id, note, tags, favorite, state, updated_at, version) VALUES ('MIGRATION_TEST', 'preserve research', '[]', false, 'inbox', now(), 7);"
)
)
command.upgrade(config, "head")
command.check(config)
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
("preserve research", 7)
]
assert asyncio.run(sql("SELECT count(*) FROM catalog_batches")) == [(0,)]
command.downgrade(config, "0002")
command.upgrade(config, "head")
command.check(config)
assert asyncio.run(sql("SELECT note, version FROM research WHERE alpha_id='MIGRATION_TEST'")) == [
("preserve research", 7)
]
print(
"PostgreSQL 17: 0002 → 0003, downgrade/re-upgrade, metadata check, Alpha/research preservation passed"
)
async def flow():
import httpx
from app.config import Settings
from app.main import create_app
from app.worldquant import WqClient
from tests.catalog_fake import catalog_response
from tests.test_catalog import SCOPE, prepare, search, sync
def upstream(request):
if request.url.path == "/authentication":
return httpx.Response(201, json={"token": {"expiry": 14400}})
if request.url.path == "/users/self":
return httpx.Response(200, json={"id": "PG_TEST_USER"})
assert request.method == "GET"
return catalog_response(request) or httpx.Response(404)
settings = Settings(_env_file=None, enable_runner=False, public_origin="http://testserver")
app = create_app(settings, WqClient(settings, transport=httpx.MockTransport(upstream)))
async with app.router.lifespan_context(app):
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
headers={"X-WQ-Request": "1"},
) as client:
assert (
await client.post(
"/api/v1/auth/login", json={"username": "admin", "password": "migration-test-only"}
)
).status_code == 200
await client.put(
"/api/v1/account/credentials",
json={"email": "pg@example.com", "password": "synthetic-only"},
)
job = (await client.post("/api/v1/account/connect")).json()
await app.state.runner.execute(job["id"])
catalog = (client, app.state.runner, {})
assert (await sync(catalog))["status"] == "completed"
version = (await sync(catalog, "TEST_FIN"))["id"]
result = await search(client, "/datasets/TEST_FIN/fields")
assert result["complete_count"] == 123
draft = (await prepare(client, version)).json()
assert len(draft["field_ids"]) == 123
responses = await asyncio.gather(
*[
client.patch(
"/api/v1/catalog/datasets/TEST_FIN/research",
params=SCOPE,
json={"version": 1, "note": value},
)
for value in ["one", "two"]
]
)
assert sorted(r.status_code for r in responses) == [200, 409]
await sync(catalog, "TEST_FIN")
assert (await prepare(client, version)).status_code == 409
persisted = (await client.get("/api/v1/catalog/inputs/" + draft["id"])).json()
assert persisted == draft
print(
"PostgreSQL: real API/runner multi-page dedupe, immutable draft, refresh conflict and concurrent note CAS passed"
)
asyncio.run(flow())
+257
View File
@@ -0,0 +1,257 @@
"""Public API through real business/runner/database; only upstream HTTP is replaced."""
import asyncio
import httpx
import pytest
from app.jobs import Runner
from app.models import Job
from app.worldquant import WqClient
from tests.catalog_fake import catalog_response, field_records
SCOPE = dict(instrument_type="EQUITY", region="USA", universe="TOP3000", delay=1)
BASE = "/api/v1/catalog"
@pytest.fixture
async def catalog(logged_in, app):
state = {"fail": False, "fields": field_records(), "calls": [], "mode": "", "block": None}
async def upstream(request):
state["calls"].append((request.url.path, int(request.url.params.get("offset", 0))))
if request.url.path == "/authentication":
if state.get("persona"):
return httpx.Response(
401, headers={"WWW-Authenticate": "persona", "Location": "/authentication/persona/test"}
)
return httpx.Response(201, json={"token": {"expiry": 14400}})
assert request.method == "GET"
if request.url.path == "/users/self":
return httpx.Response(200, json={"id": "TEST_USER"})
if request.url.path == "/data-fields":
if state.get("throttle"):
state["throttle"] = False
return httpx.Response(429, headers={"Retry-After": "2"})
if state["mode"] == "invalid-next":
return httpx.Response(200, json={"results": state["fields"][:50], "next": []})
if state["mode"] == "missing-owner":
return httpx.Response(200, json={"results": [{"id": "UNOWNED"}], "next": None})
if state["mode"] == "coverage-unit":
return httpx.Response(
200, json={"results": [{**state["fields"][0], "coverage": 95}], "next": None}
)
if int(request.url.params["offset"]) >= 50:
if state["block"]:
state["block"].set()
await asyncio.Future()
if state["fail"]:
return httpx.Response(403)
if state["mode"] == "early":
return httpx.Response(200, json={"results": [], "next": "/next", "count": 123})
if state["mode"] == "repeat":
return httpx.Response(200, json={"results": state["fields"][:50], "next": "/next"})
if state["mode"] == "wrong-owner":
return httpx.Response(200, json={"results": field_records("OTHER", 1)})
return catalog_response(request, state["fields"]) or httpx.Response(404)
await app.state.runner.client.close()
app.state.runner.client = WqClient(app.state.settings, transport=httpx.MockTransport(upstream))
client = logged_in
assert (
await client.put(
"/api/v1/account/credentials", json={"email": "test@example.com", "password": "test-only"}
)
).status_code == 200
connect = (await client.post("/api/v1/account/connect")).json()
await app.state.runner.execute(connect["id"])
return client, app.state.runner, state
async def sync(catalog, dataset=None, scope=SCOPE):
client, runner, _ = catalog
response = await client.post(BASE + "/sync-jobs", json={"scope": scope, "dataset_id": dataset})
assert response.status_code == 202, response.text
job = response.json()
await runner.execute(job["id"])
return (await client.get("/api/v1/sync-jobs/" + job["id"])).json()
async def search(client, suffix="/datasets", **params):
response = await client.get(BASE + suffix, params={**SCOPE, **params})
assert response.status_code == 200, response.text
return response.json()
async def prepare(client, version, **changes):
return await client.post(
BASE + "/inputs",
json={
"scope": SCOPE,
"dataset_id": "TEST_FIN",
"collection_version": version,
"selection": "all",
**changes,
},
)
async def test_complete_workflow_filters_notes_immutable_input(catalog):
client, _, state = catalog
assert (await search(client))["total"] == 0
assert (await sync(catalog))["status"] == "completed"
datasets = await search(client, category="基本面", subcategory="财务报表")
assert [r["id"] for r in datasets["items"]] == ["TEST_FIN"]
assert datasets["items"][0]["complete_count"] is None
assert (await sync(catalog, "TEST_FIN"))["processed"] == 123
fields = await search(client, "/datasets/TEST_FIN/fields", q="字段 12", limit=1)
assert fields["total"] == 3 and fields["complete_count"] == 123 and len(fields["items"]) == 1
version = fields["collection_version"]
response = await prepare(client, version)
assert response.status_code == 201, response.text
draft = response.json()
assert len(draft["field_ids"]) == 123 and draft["status"] == "draft"
for suffix in ["/datasets/TEST_FIN", "/datasets/TEST_FIN/fields/TEST_FIN_122"]:
detail = await search(client, suffix)
assert detail["research"]["version"] == 1
response = await client.patch(
BASE + suffix + "/research", params=SCOPE, json={"version": 1, "note": "保留研究备注"}
)
assert response.status_code == 200
assert (
await client.patch(
BASE + suffix + "/research", params=SCOPE, json={"version": 1, "note": "不能覆盖"}
)
).status_code == 409
detail = await search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122")
assert detail["coverage"] is None and detail["unit"] is None and detail["field_type"] == "FUTURE_TYPE"
assert (await search(client, "/datasets/TEST_FIN/fields", coverage_min=0))["total"] == 122
state["fields"] = field_records(count=125)
assert (await sync(catalog, "TEST_FIN"))["processed"] == 125
assert (await sync(catalog))["status"] == "completed"
newer = await search(client, "/datasets/TEST_FIN/fields")
assert newer["collection_version"] != version
assert (await prepare(client, version)).status_code == 409
assert len((await prepare(client, newer["collection_version"])).json()["field_ids"]) == 125
assert (await client.get(BASE + "/inputs/" + draft["id"])).json() == draft
assert (await search(client, "/datasets/TEST_FIN/fields/TEST_FIN_122"))["research"][
"note"
] == "保留研究备注"
assert (await search(client, "/datasets/TEST_FIN"))["research"]["note"] == "保留研究备注"
async def test_partial_refresh_resume_cancel_restart_keeps_old_version(catalog):
client, runner, state = catalog
await sync(catalog)
state["fail"] = True
job = await sync(catalog, "TEST_FIN")
assert job["status"] == "failed" and job["processed"] == 50
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] is None
assert (await prepare(client, job["id"])).status_code == 409
state["fail"] = False
state["calls"].clear()
assert (await client.post("/api/v1/sync-jobs/" + job["id"] + "/retry")).status_code == 200
await runner.execute(job["id"])
assert state["calls"][0] == ("/data-fields", 50)
old_version = (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"]
state["fail"] = True
refresh = await sync(catalog, "TEST_FIN")
assert refresh["status"] == "failed"
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] == old_version
state["fail"] = False
state["calls"].clear()
async with runner.sessions() as db:
row = await db.get(Job, refresh["id"])
row.status = "running"
await db.commit()
restarted = Runner(runner.sessions, runner.settings, runner.client)
await restarted.start()
async with asyncio.timeout(5):
while True:
response = (await client.get("/api/v1/sync-jobs/" + refresh["id"])).json()
if response["status"] in ("completed", "failed"):
break
await asyncio.sleep(0.02)
assert response["status"] == "completed"
assert state["calls"][0] == ("/data-fields", 50)
state["block"] = asyncio.Event()
response = await client.post(BASE + "/sync-jobs", json={"scope": SCOPE, "dataset_id": "TEST_FIN"})
cancel_id = response.json()["id"]
restarted.wake.set()
await asyncio.wait_for(state["block"].wait(), 5)
await client.post("/api/v1/sync-jobs/" + cancel_id + "/cancel")
await restarted.cancel(cancel_id)
assert (await client.get("/api/v1/sync-jobs/" + cancel_id)).json()["status"] == "cancelled"
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] == refresh["id"]
await restarted.stop()
@pytest.mark.parametrize(
"mode", ["early", "repeat", "wrong-owner", "invalid-next", "missing-owner", "coverage-unit"]
)
async def test_anomalous_pagination_is_never_complete(catalog, mode):
client, _, state = catalog
await sync(catalog)
state["mode"] = mode
assert (await sync(catalog, "TEST_FIN"))["status"] == "failed"
assert (await search(client, "/datasets/TEST_FIN/fields"))["collection_version"] is None
async def test_scope_ownership_empty_and_unknown_fields_are_rejected(catalog):
client, _, _ = catalog
await sync(catalog)
version = (await sync(catalog, "TEST_FIN"))["id"]
assert (await prepare(client, version, selection="explicit", excluded_ids=["OTHER"])).status_code == 422
assert (
await prepare(
client, version, selection="explicit", excluded_ids=[f"TEST_FIN_{i:03}" for i in range(123)]
)
).status_code == 422
assert (await prepare(client, version, dataset_id="TEST_NEWS")).status_code == 409
assert (await prepare(client, version, scope={**SCOPE, "delay": 0})).status_code == 404
assert (await prepare(client, version, scope={**SCOPE, "region": "CHN"})).status_code == 422
subset = await prepare(client, version, selection="explicit", excluded_ids=["TEST_FIN_110"])
assert subset.status_code == 201 and len(subset.json()["field_ids"]) == 122
assert "TEST_FIN_110" not in subset.json()["field_ids"]
other = {**SCOPE, "delay": 0}
await sync(catalog, scope=other)
await sync(catalog, "TEST_FIN", scope=other)
assert (await prepare(client, version, scope=other)).status_code == 409
assert len((await client.get(BASE + "/inputs", params=SCOPE)).json()) == 1
async def test_catalog_authentication_and_origin(app, client):
assert (await client.get(BASE + "/datasets", params=SCOPE)).status_code == 401
assert (
await client.post(BASE + "/sync-jobs", headers={"Origin": "http://evil.test"}, json={"scope": SCOPE})
).status_code == 403
async def test_retry_after_auth_wait_disconnect_and_collection_manifest(catalog):
client, runner, state = catalog
await sync(catalog)
delays = []
async def sleep(delay):
delays.append(delay)
runner.client.sleep = sleep
state["throttle"] = True
job = await sync(catalog, "TEST_FIN")
assert job["status"] == "completed" and delays == [2]
manifest = await search(client, "/datasets/TEST_FIN/collection")
assert manifest["collection_version"] == job["id"] and len(manifest["field_ids"]) == 123
state["persona"] = True
runner.client.authenticated = False
waiting = await sync(catalog, "TEST_FIN")
assert waiting["status"] == "waiting_auth"
assert (await search(client, "/datasets/TEST_FIN/collection")) == manifest
await runner.disconnect()
assert (await client.get("/api/v1/sync-jobs/" + waiting["id"])).json()["status"] == "waiting_connection"
assert (await client.post(BASE + "/sync-jobs", json={"scope": SCOPE})).status_code == 409
# Explicit reconnect verifies the original account and resumes the same task.
state["persona"] = False
connect = (await client.post("/api/v1/account/connect")).json()
await runner.execute(connect["id"])
await runner.execute(waiting["id"])
assert (await client.get("/api/v1/sync-jobs/" + waiting["id"])).json()["status"] == "completed"