"""REST API 路由(全部挂在 /api 下)。看板只读自有库。""" from __future__ import annotations import logging from datetime import date, datetime, timedelta from typing import Dict, List, Optional from fastapi import APIRouter, Body, Depends, HTTPException, Query, UploadFile from sqlalchemy import delete, func, select from sqlalchemy.orm import Session from .config import INDEX_NAMES, get_settings from .db import ( CycleAnnotation, EtfFlow, IndexDaily, MarketEvent, SentimentDaily, engine, get_session, ) from .importers import parse_etf_csv, parse_events_csv, list_auto_sources from .schemas import ( AnnotationIn, EtfFlowIn, EtfFlowOut, EventIn, EventOut, ImportResult, ) from .signals import compute_signals log = logging.getLogger("as-event.api") settings = get_settings() router = APIRouter(prefix="/api") # ============================ 工具 ============================ def _parse_qdate(s: Optional[str], default: date) -> date: if not s: return default for fmt in ("%Y%m%d", "%Y-%m-%d"): try: return datetime.strptime(s, fmt).date() except ValueError: continue raise HTTPException(400, f"日期格式错误: {s}") def _default_range() -> tuple[date, date]: end = date.today() start = end - timedelta(days=730) # 默认近两年 return start, end # ============================ 健康 ============================ @router.get("/health") def health() -> dict: own_ok = False try: with engine.connect() as c: c.execute(select(func.count()).select_from(IndexDaily)) own_ok = True except Exception as e: # noqa: BLE001 log.warning("own db 健康检查失败: %s", e) return { "status": "ok" if own_ok else "degraded", "own_db": own_ok, "legacy_mysql_configured": bool(settings.legacy_mysql_dsn), "legacy_pg_configured": bool(settings.legacy_pg_dsn), "ohlc_provider": settings.ohlc_provider, "auto_sources": list_auto_sources(), } # ============================ 指数 / K线 ============================ @router.get("/index/list") def index_list(session: Session = Depends(get_session)) -> List[dict]: out = [] for code in settings.index_code_list: rng = session.execute( select(func.min(IndexDaily.trade_date), func.max(IndexDaily.trade_date)) .where(IndexDaily.index_code == code) ).one() out.append( { "code": code, "name": INDEX_NAMES.get(code, code), "data_start": rng[0].isoformat() if rng[0] else None, "data_end": rng[1].isoformat() if rng[1] else None, } ) return out @router.get("/index/{code}/kline") def index_kline( code: str, start: Optional[str] = Query(None), end: Optional[str] = Query(None), session: Session = Depends(get_session), ) -> dict: ds, de = _default_range() ds = _parse_qdate(start, ds) de = _parse_qdate(end, de) # 为滚动分位预留前置缓冲,保证请求区间起点温度也有意义 buffer_start = ds - timedelta(days=int(settings.vol_window * 2)) rows = ( session.execute( select(IndexDaily) .where(IndexDaily.index_code == code, IndexDaily.trade_date >= buffer_start, IndexDaily.trade_date <= de) .order_by(IndexDaily.trade_date.asc()) ) .scalars() .all() ) if not rows: return {"index_code": code, "index_name": INDEX_NAMES.get(code, code), "ohlc_available": False, "dates": [], "ohlc": [], "amount": [], "volume": [], "pct_chg": [], "temperature": [], "flags": []} bars = [{"trade_date": r.trade_date, "close": _f(r.close), "amount": _f(r.amount)} for r in rows] # 情绪、ETF 用于温度计算 senti = _sentiment_map(session, buffer_start, de) etf_net = _etf_net_map(session, buffer_start, de, related_index=code) sig = compute_signals(bars, senti, etf_net) sig_by_date = {s["trade_date"]: s for s in sig} # 只输出请求区间 dates, ohlc, amount, volume, pct, temp, flags = [], [], [], [], [], [], [] ohlc_available = False for r in rows: if r.trade_date < ds: continue dates.append(r.trade_date.isoformat()) o, h, l, c = _f(r.open), _f(r.high), _f(r.low), _f(r.close) ohlc.append([o, h, l, c]) if r.ohlc_source and r.ohlc_source != "close_only": ohlc_available = True amount.append(_f(r.amount)) volume.append(r.volume) pct.append(_f(r.pct_chg)) s = sig_by_date.get(r.trade_date, {}) temp.append(s.get("temperature")) flags.append(s.get("flags", [])) return { "index_code": code, "index_name": INDEX_NAMES.get(code, code), "ohlc_available": ohlc_available, # False → 前端用收盘线 "dates": dates, "ohlc": ohlc, "amount": amount, "volume": volume, "pct_chg": pct, "temperature": temp, "flags": flags, "thresholds": {"hot": settings.hot_threshold, "cold": settings.cold_threshold}, } # ============================ 情绪 ============================ @router.get("/sentiment") def sentiment( start: Optional[str] = Query(None), end: Optional[str] = Query(None), session: Session = Depends(get_session), ) -> List[dict]: ds, de = _default_range() ds, de = _parse_qdate(start, ds), _parse_qdate(end, de) rows = ( session.execute( select(SentimentDaily) .where(SentimentDaily.trade_date >= ds, SentimentDaily.trade_date <= de) .order_by(SentimentDaily.trade_date.asc()) ).scalars().all() ) return [ { "trade_date": r.trade_date.isoformat(), "up_down_ratio": _f(r.up_down_ratio), "median_pct_chg": _f(r.median_pct_chg), "pct_chg_gt_5_count": r.pct_chg_gt_5_count, "limit_up_count": r.limit_up_count, "limit_down_count": r.limit_down_count, "margin_balance": _f(r.margin_balance), "new_accounts": r.new_accounts, } for r in rows ] # ============================ ETF ============================ @router.get("/etf") def etf_list( start: Optional[str] = Query(None), end: Optional[str] = Query(None), index: Optional[str] = Query(None), session: Session = Depends(get_session), ) -> List[EtfFlowOut]: ds, de = _default_range() ds, de = _parse_qdate(start, ds), _parse_qdate(end, de) q = select(EtfFlow).where(EtfFlow.trade_date >= ds, EtfFlow.trade_date <= de) if index: q = q.where(EtfFlow.related_index == index) rows = session.execute(q.order_by(EtfFlow.trade_date.asc())).scalars().all() return [EtfFlowOut.model_validate(r) for r in rows] @router.post("/etf", response_model=EtfFlowOut) def etf_create(payload: EtfFlowIn, session: Session = Depends(get_session)) -> EtfFlowOut: obj = EtfFlow(**payload.model_dump(), source="manual") session.add(obj) session.commit() session.refresh(obj) return EtfFlowOut.model_validate(obj) @router.post("/etf/import", response_model=ImportResult) async def etf_import(file: UploadFile, session: Session = Depends(get_session)) -> ImportResult: content = (await file.read()).decode("utf-8-sig") try: recs = parse_etf_csv(content) except ValueError as e: raise HTTPException(400, str(e)) ins = 0 for rec in recs: session.add(EtfFlow(**rec)) ins += 1 session.commit() return ImportResult(inserted=ins, updated=0, total=len(recs)) # ============================ 事件 CRUD ============================ @router.get("/events") def events_list( start: Optional[str] = Query(None), end: Optional[str] = Query(None), index: Optional[str] = Query(None), category: Optional[str] = Query(None), session: Session = Depends(get_session), ) -> List[EventOut]: ds, de = _default_range() ds, de = _parse_qdate(start, ds), _parse_qdate(end, de) q = select(MarketEvent).where(MarketEvent.event_date >= ds, MarketEvent.event_date <= de) if category: q = q.where(MarketEvent.category == category) rows = session.execute(q.order_by(MarketEvent.event_date.asc())).scalars().all() # related_indices 过滤(CSV 包含匹配,空视为全市场) if index: rows = [r for r in rows if (not r.related_indices) or (index in r.related_indices)] return [EventOut.model_validate(r) for r in rows] @router.post("/events", response_model=EventOut) def event_create(payload: EventIn, session: Session = Depends(get_session)) -> EventOut: obj = MarketEvent(**payload.model_dump(), source="manual") session.add(obj) session.commit() session.refresh(obj) return EventOut.model_validate(obj) @router.put("/events/{event_id}", response_model=EventOut) def event_update(event_id: int, payload: EventIn, session: Session = Depends(get_session)) -> EventOut: obj = session.get(MarketEvent, event_id) if not obj: raise HTTPException(404, "事件不存在") for k, v in payload.model_dump().items(): setattr(obj, k, v) session.commit() session.refresh(obj) return EventOut.model_validate(obj) @router.delete("/events/{event_id}") def event_delete(event_id: int, session: Session = Depends(get_session)) -> dict: obj = session.get(MarketEvent, event_id) if not obj: raise HTTPException(404, "事件不存在") session.delete(obj) session.commit() return {"deleted": event_id} @router.post("/events/import", response_model=ImportResult) async def events_import(file: UploadFile, session: Session = Depends(get_session)) -> ImportResult: content = (await file.read()).decode("utf-8-sig") try: evs = parse_events_csv(content) except ValueError as e: raise HTTPException(400, str(e)) for ev in evs: session.add(MarketEvent(**ev)) session.commit() return ImportResult(inserted=len(evs), updated=0, total=len(evs)) # ============================ 人工顶底标注 ============================ @router.get("/annotations") def anno_list(index: Optional[str] = Query(None), session: Session = Depends(get_session)) -> List[dict]: q = select(CycleAnnotation) if index: q = q.where(CycleAnnotation.index_code == index) rows = session.execute(q.order_by(CycleAnnotation.anno_date.asc())).scalars().all() return [ {"id": r.id, "index_code": r.index_code, "anno_date": r.anno_date.isoformat(), "kind": r.kind, "note": r.note} for r in rows ] @router.post("/annotations") def anno_create(payload: AnnotationIn, session: Session = Depends(get_session)) -> dict: obj = CycleAnnotation(**payload.model_dump()) session.add(obj) session.commit() session.refresh(obj) return {"id": obj.id} @router.delete("/annotations/{anno_id}") def anno_delete(anno_id: int, session: Session = Depends(get_session)) -> dict: session.execute(delete(CycleAnnotation).where(CycleAnnotation.id == anno_id)) session.commit() return {"deleted": anno_id} # ============================ 触发 ETL ============================ @router.post("/sync/index") def sync_index_ep( start: str = Body(..., embed=True), end: Optional[str] = Body(None, embed=True), codes: Optional[str] = Body(None, embed=True), ) -> dict: from .etl import sync_index, _parse_date # 延迟导入避免循环 de = _parse_date(end) if end else date.today() code_list = codes.split(",") if codes else None res = sync_index(_parse_date(start), de, code_list) return {"synced": res} @router.post("/sync/sentiment") def sync_sentiment_ep( start: str = Body(..., embed=True), end: Optional[str] = Body(None, embed=True), ) -> dict: from .etl import sync_sentiment, _parse_date de = _parse_date(end) if end else date.today() n = sync_sentiment(_parse_date(start), de) return {"synced": n} # ============================ helpers ============================ def _f(v) -> Optional[float]: return float(v) if v is not None else None def _sentiment_map(session: Session, ds: date, de: date) -> Dict[date, dict]: rows = session.execute( select(SentimentDaily).where(SentimentDaily.trade_date >= ds, SentimentDaily.trade_date <= de) ).scalars().all() return { r.trade_date: { "up_down_ratio": _f(r.up_down_ratio), "pct_chg_gt_5_count": r.pct_chg_gt_5_count, } for r in rows } def _etf_net_map(session: Session, ds: date, de: date, related_index: Optional[str] = None) -> Dict[date, float]: q = select(EtfFlow.trade_date, func.sum(EtfFlow.net_inflow)).where( EtfFlow.trade_date >= ds, EtfFlow.trade_date <= de ) # 温度用全市场 ETF 净流入;如需按指数可加过滤,这里保留全市场以反映整体资金 q = q.group_by(EtfFlow.trade_date) rows = session.execute(q).all() return {d: float(s) for d, s in rows if s is not None}