"""Publish complete enumerations only; retain staging checkpoints and old versions.""" import asyncio import math import re from urllib.parse import parse_qs, urlparse from uuid import uuid4 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, value_score=number(raw.get("valueScore")) if not dataset_id else None, 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, *, 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: 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) 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, batch_id) added = 0 for entry in entries: if await db.get(CatalogEntry, (batch_id, entry["id"])): continue db.add(CatalogEntry(batch_id=batch_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 if not full: job.processed = batch.count offset += len(rows) 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() 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 = batch_id else: scope_row = await db.get(CatalogScope, scope.key()) scope_row.catalog_version, scope_row.synced_at = batch_id, now() ids = ( 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)): db.add(CatalogDataset(scope_key=scope.key(), id=item_id)) 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