383 lines
13 KiB
Python
383 lines
13 KiB
Python
"""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}
|