as-event/backend/app/api.py

383 lines
13 KiB
Python
Raw Permalink Normal View History

2026-07-27 09:05:52 +08:00
"""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}