"""数据同步 — 接收 local 端推送的 Parquet 数据包。""" from __future__ import annotations import logging import tarfile import tempfile from io import BytesIO from pathlib import Path 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 _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 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 target.write_bytes(src.read()) extracted += 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", parts, extracted) return {"ok": True, "parts": sorted(requested), "file_count": extracted}