"""数据同步 — 接收 local 端推送的 Parquet 数据包。""" from __future__ import annotations import logging import os import tarfile import tempfile from io import BytesIO from pathlib import Path import polars as pl from fastapi import APIRouter, Depends, HTTPException, Request, UploadFile, File, Form from app.config import settings logger = logging.getLogger(__name__) router = APIRouter(prefix="/api/data/sync", tags=["data-sync"]) # 允许同步的子目录列表 SYNCABLE_PARTS = { "kline_daily", "kline_daily_enriched", "kline_index_daily", "kline_index_enriched", "kline_etf_daily", "kline_etf_enriched", "kline_etf_minute", "kline_minute", "adj_factor", "adj_factor_etf", "instruments", "instruments_index", "instruments_etf", "financials", "pools", "ext_data", } def _resolve_unique_key(columns: list[str]) -> list[str] | None: """根据列名推断去重键。 策略: - 含 symbol + datetime: 按 [symbol, datetime] 去重(分钟 K) - 含 symbol + date: 按 [symbol, date] 去重(日 K / 复权因子) - 含 symbol: 按 [symbol] 去重(标的维表) - 其他: 无法推断,回退到覆盖 """ cols = set(columns) if "symbol" not in cols: return None if "datetime" in cols: return ["symbol", "datetime"] if "date" in cols: return ["symbol", "date"] return ["symbol"] def _merge_parquet(target: Path, new_bytes: bytes) -> tuple[bool, str]: """把 new_bytes 代表的 Parquet 与 target 已有文件合并去重。 返回 (是否成功落盘, 操作描述): - created: 目标不存在,直接新建 - merged: 与已有文件按主键合并去重(保留后写入的记录) - empty_new: 上传文件为空,保留旧文件 - overwritten_no_key: 无法推断去重键,回退覆盖 - overwritten_corrupt_existing: 已有文件损坏,回退覆盖 - invalid_new: 上传文件不是合法 Parquet,未落盘 """ if not target.exists(): target.parent.mkdir(parents=True, exist_ok=True) target.write_bytes(new_bytes) return True, "created" try: new_df = pl.read_parquet(BytesIO(new_bytes)) except Exception as exc: logger.warning("failed to parse uploaded parquet %s: %s", target.name, exc) return False, "invalid_new" if new_df.is_empty(): return True, "empty_new" key = _resolve_unique_key(new_df.columns) if key is None: logger.warning( "no unique key for %s (columns=%s), overwriting", target.name, new_df.columns, ) target.write_bytes(new_bytes) return True, "overwritten_no_key" try: existing_df = pl.read_parquet(target) except Exception as exc: logger.warning("existing parquet %s corrupt, overwriting: %s", target, exc) target.write_bytes(new_bytes) return True, "overwritten_corrupt_existing" if existing_df.is_empty(): merged = new_df action = "created" else: merged = pl.concat([existing_df, new_df], how="diagonal_relaxed").unique( subset=key, keep="last" ) action = "merged" # 原子写入,防止写入过程中崩溃导致文件损坏 tmp = target.with_suffix(f".tmp-{os.getpid()}") try: merged.write_parquet(tmp) os.replace(tmp, target) except Exception: if tmp.exists(): tmp.unlink(missing_ok=True) raise return True, action def _verify_sync_key(request: Request) -> None: """简单的鉴权 — 校验 X-Sync-Key 头。""" key = request.headers.get("X-Sync-Key", "") expected = getattr(settings, "sync_key", "") if expected and key != expected: raise HTTPException(status_code=403, detail="sync key mismatch") def _refresh_views(request: Request) -> None: """重新注册 DuckDB 视图 + 刷新 Polars 缓存。""" repo = getattr(request.app.state, "repo", None) if repo is None: return try: d = repo.store.data_dir.as_posix() views = { "kline_daily": f"{d}/kline_daily/**/*.parquet", "kline_enriched": f"{d}/kline_daily_enriched/**/*.parquet", "kline_index_daily": f"{d}/kline_index_daily/**/*.parquet", "kline_index_enriched": f"{d}/kline_index_enriched/**/*.parquet", "kline_etf_daily": f"{d}/kline_etf_daily/**/*.parquet", "kline_etf_enriched": f"{d}/kline_etf_enriched/**/*.parquet", "kline_etf_minute": f"{d}/kline_etf_minute/**/*.parquet", "kline_minute": f"{d}/kline_minute/**/*.parquet", "adj_factor": f"{d}/adj_factor/**/*.parquet", "adj_factor_etf": f"{d}/adj_factor_etf/**/*.parquet", "instruments": f"{d}/instruments/**/*.parquet", "instruments_index": f"{d}/instruments_index/**/*.parquet", "instruments_etf": f"{d}/instruments_etf/**/*.parquet", } for name, path in views.items(): try: repo.store.db.execute( f"CREATE OR REPLACE VIEW {name} AS " f"SELECT * FROM read_parquet('{path}', union_by_name=true)" ) except Exception as exc: logger.warning("view %s refresh failed: %s", name, exc) repo.store._register_unified_views() repo.clear_cache() repo.refresh_cache() # 清除 API 数据缓存 from app.api.data import invalidate_data_cache invalidate_data_cache() except Exception as exc: logger.warning("cache refresh failed: %s", exc) @router.post("/upload") async def upload_sync( request: Request, file: UploadFile = File(...), parts: str = Form(""), _auth: None = Depends(_verify_sync_key), ): """接收 local 端打包的 tar.gz 数据,解压写入 data/ 目录。""" data_dir: Path = request.app.state.datastore.data_dir # 解析要同步的 parts requested = set(p.strip() for p in parts.split(",") if p.strip()) if not requested: raise HTTPException(status_code=400, detail="parts 不能为空,逗号分隔子目录名") invalid = requested - SYNCABLE_PARTS if invalid: raise HTTPException( status_code=400, detail=f"不支持的同步目录: {', '.join(sorted(invalid))}", ) # 读取上传的 tar.gz raw = await file.read() if not raw: raise HTTPException(status_code=400, detail="上传文件为空") extracted = 0 action_counter: dict[str, int] = {} try: with tarfile.open(fileobj=BytesIO(raw), mode="r:gz") as tar: for member in tar.getmembers(): # 安全校验: 防止路径穿越 member_path = Path(member.name) if member_path.is_absolute() or ".." in member_path.parts: logger.warning("skipping unsafe path: %s", member.name) continue # 确定属于哪个 part top_dir = member_path.parts[0] if member_path.parts else "" if top_dir not in requested: continue target = data_dir / member.name if member.isfile(): target.parent.mkdir(parents=True, exist_ok=True) with tar.extractfile(member) as src: if src is None: continue ok, action = _merge_parquet(target, src.read()) if ok: extracted += 1 action_counter[action] = action_counter.get(action, 0) + 1 except tarfile.TarError as exc: raise HTTPException(status_code=400, detail=f"压缩包解析失败: {exc}") from exc # 刷新缓存 _refresh_views(request) logger.info("sync uploaded: parts=%s files=%d actions=%s", parts, extracted, action_counter) return {"ok": True, "parts": sorted(requested), "file_count": extracted, "actions": action_counter}