功能更新

This commit is contained in:
2026-08-15 08:57:15 +08:00
parent 50fd032b45
commit c1c43d2ff7
30 changed files with 1908 additions and 888 deletions

View File

@@ -2,6 +2,7 @@
GET /api/health 健康检查
GET /api/candles/{sym} 取 K 线(支持 1d/1w/1M/1y 周期,日线为基底聚合)
GET /api/stocks 全市场股票列表(基本信息 + 最新行情 + 缓存条数)
POST /api/backtest 跑回测,返回 K线+指标+买卖点+净值+绩效
POST /api/screener/run 智能选股:自然语言 -> 条件 -> 全市场筛选
POST /api/screener/sync 启动全市场数据同步(后台任务)
@@ -9,45 +10,66 @@
"""
from __future__ import annotations
import bisect
import json
import pandas as pd
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy import select
from sqlalchemy import select, text
from sqlalchemy.ext.asyncio import AsyncSession
from .backtest.engine import BacktestConfig, run_backtest
from .auth import require_user
from .backtest.events import EventEngineError, run_event_backtest
from .backtest.strategies import build_strategy
from .config import settings
from .data import fetcher, repository
from .data.aggregation import bars_per_year, resample_bars
from .data.symbols import plain_code
from .data.synthetic import seed_if_empty
from .db import get_session
from .domain import Bar
from . import indicators as ind
from .models import BacktestRun, DailySnapshot, MarketDaily, StockBasic
from .models import (
AdjFactor,
BacktestRun,
DailySnapshot,
MarketDaily,
ScreenerQuery,
StockBasic,
UserPreference,
WatchlistItem,
)
from .schemas import (
BacktestRequest,
BacktestResponse,
CandleOut,
EquityPoint,
EventBacktestRequest,
EventBacktestResponse,
IndicatorOut,
MetricsOut,
PreferencesOut,
PreferencesUpdate,
PreviewInfoOut,
PreviewResponse,
ScreenerQueryListResponse,
ScreenerQueryOut,
ScreenerRunRequest,
ScreenerRunResponse,
ScreenerSyncRequest,
ScreenerSyncStatus,
SignalOut,
StockListItemOut,
StockListResponse,
StockFacetsResponse,
FacetItemOut,
SyncRequest,
SyncResponse,
WatchlistOp,
)
from .screener import engine, market_sync
from .screener.engine import DataNotReadyError
from .screener.llm import ScreenerError, parse_conditions
from .screener.llm import ScreenerError, parse_conditions, parse_event_spec
router = APIRouter(prefix="/api", dependencies=[Depends(require_user)])
@@ -67,6 +89,41 @@ def _rows_to_bars(rows) -> list[Bar]:
return [Bar(ts=r.ts, open=r.open, high=r.high, low=r.low, close=r.close, volume=r.volume) for r in rows]
_ADJUST_MODES = ("bfq", "qfq", "hfq")
def _adjust_bars(bars: list[Bar], factors, from_mode: str, to_mode: str) -> list[Bar]:
"""按复权因子把 K 线从 from_mode 换算到 to_modebfq/qfq/hfq
相对不复权的乘数bfq=1qfq=f(t)/f(latest)hfq=f(t)。
因子缺失的日期向前沿用最近因子(因子是阶梯函数,除权日之间不变)。
"""
fd = sorted((f.trade_date.date(), float(f.adj_factor)) for f in factors)
fdates = [d for d, _ in fd]
f_latest = fd[-1][1]
def _f_at(d) -> float:
i = bisect.bisect_right(fdates, d) - 1
return fd[i][1] if i >= 0 else fd[0][1]
def _mult(mode: str, f: float) -> float:
if mode == "bfq":
return 1.0
return f / f_latest if mode == "qfq" else f
out: list[Bar] = []
for b in bars:
f = _f_at(b.ts.date())
m = _mult(to_mode, f) / _mult(from_mode, f)
out.append(Bar(
ts=b.ts,
open=round(b.open * m, 3), high=round(b.high * m, 3),
low=round(b.low * m, 3), close=round(b.close * m, 3),
volume=b.volume,
))
return out
@router.get("/candles/{symbol}", response_model=list[CandleOut])
async def get_candles(
symbol: str,
@@ -74,7 +131,6 @@ async def get_candles(
limit: int = 5000,
session: AsyncSession = Depends(get_session),
) -> list[CandleOut]:
await seed_if_empty(session, symbol="DEMO")
# 始终以日线为基底,再聚合到目标周期
rows = await repository.get_candles(session, symbol, "1d", limit=limit)
bars = resample_bars(_rows_to_bars(rows), timeframe)
@@ -93,15 +149,114 @@ async def sync_data(req: SyncRequest, session: AsyncSession = Depends(get_sessio
raise HTTPException(status_code=502, detail=str(e))
# ---------- 股票列表(全市场浏览) ----------
_STOCKS_SQL = text("""
SELECT sb.ts_code, sb.symbol, sb.name, sb.industry, sb.market,
c.close AS close, p.close AS prev_close, c.ts AS last_ts, cnt.n AS bar_count,
CASE WHEN c.close IS NOT NULL AND p.close IS NOT NULL AND p.close <> 0
THEN round(((c.close / p.close - 1) * 100)::numeric, 2) END AS pct_chg,
(w.id IS NOT NULL) AS watched
FROM stock_basic sb
LEFT JOIN LATERAL (
SELECT close, ts FROM candles
WHERE symbol = sb.symbol AND timeframe = '1d'
ORDER BY ts DESC LIMIT 1
) c ON true
LEFT JOIN LATERAL (
SELECT close FROM candles
WHERE symbol = sb.symbol AND timeframe = '1d' AND ts < c.ts
ORDER BY ts DESC LIMIT 1
) p ON c.ts IS NOT NULL
LEFT JOIN LATERAL (
SELECT count(*) AS n FROM candles
WHERE symbol = sb.symbol AND timeframe = '1d'
) cnt ON true
LEFT JOIN watchlist_items w ON w.ts_code = sb.ts_code AND w.user_id = :uid
WHERE sb.list_status = 'L'
AND (:search = '' OR sb.symbol LIKE :psearch OR sb.name LIKE :psearch)
AND (:market = '' OR sb.market = :market)
AND (:industry = '' OR sb.industry = :industry)
AND (:area = '' OR sb.area = :area)
AND (:watched_only = false OR w.id IS NOT NULL)
ORDER BY w.id DESC NULLS LAST, sb.symbol
LIMIT :limit OFFSET :offset
""")
_STOCKS_COUNT_SQL = text("""
SELECT count(*) FROM stock_basic sb
LEFT JOIN watchlist_items w ON w.ts_code = sb.ts_code AND w.user_id = :uid
WHERE sb.list_status = 'L'
AND (:search = '' OR sb.symbol LIKE :psearch OR sb.name LIKE :psearch)
AND (:market = '' OR sb.market = :market)
AND (:industry = '' OR sb.industry = :industry)
AND (:area = '' OR sb.area = :area)
AND (:watched_only = false OR w.id IS NOT NULL)
""")
@router.get("/stocks", response_model=StockListResponse)
async def list_stocks(
search: str = "",
market: str = "",
industry: str = "",
area: str = "",
watched_only: bool = False,
limit: int = 100,
offset: int = 0,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> StockListResponse:
"""全市场股票列表stock_basic 基本信息 + candles 最新行情(本地缓存,无缓存则行情列为空)。
自选股watchlist_items排最前watched_only=true 只看自选。"""
search = search.strip()
limit = max(1, min(limit, 500))
offset = max(0, offset)
params = {
"search": search,
"psearch": f"%{search}%",
"market": market,
"industry": industry,
"area": area,
"watched_only": watched_only,
"uid": user.id,
"limit": limit,
"offset": offset,
}
total = (await session.execute(_STOCKS_COUNT_SQL, params)).scalar_one()
rows = (await session.execute(_STOCKS_SQL, params)).mappings().all()
return StockListResponse(total=total, items=[StockListItemOut(**r) for r in rows])
@router.get("/stocks/facets", response_model=StockFacetsResponse)
async def stock_facets(session: AsyncSession = Depends(get_session)) -> StockFacetsResponse:
"""看股页筛选项:行业 / 地域(含数量,按数量降序)。"""
industries = (
await session.execute(text("""
SELECT industry AS name, count(*) AS n FROM stock_basic
WHERE list_status = 'L' AND industry IS NOT NULL AND industry <> ''
GROUP BY industry ORDER BY n DESC
"""))
).mappings().all()
areas = (
await session.execute(text("""
SELECT area AS name, count(*) AS n FROM stock_basic
WHERE list_status = 'L' AND area IS NOT NULL AND area <> ''
GROUP BY area ORDER BY n DESC
"""))
).mappings().all()
return StockFacetsResponse(
industries=[FacetItemOut(name=r["name"], count=r["n"]) for r in industries],
areas=[FacetItemOut(name=r["name"], count=r["n"]) for r in areas],
)
@router.post("/backtest", response_model=BacktestResponse)
async def backtest(
req: BacktestRequest,
session: AsyncSession = Depends(get_session),
) -> BacktestResponse:
await seed_if_empty(session, symbol="DEMO")
# 非演示标的:首次自动拉取真实数据并缓存
if req.symbol != "DEMO" and not await fetcher.is_cached(session, req.symbol):
# 真实数据:本地无缓存则先拉取
if not await fetcher.is_cached(session, req.symbol):
try:
await fetcher.sync_symbol(session, req.symbol, source="auto")
except Exception as e: # noqa: BLE001
@@ -179,17 +334,71 @@ async def backtest(
)
@router.post("/backtest/event", response_model=EventBacktestResponse)
async def backtest_event(
req: EventBacktestRequest,
session: AsyncSession = Depends(get_session),
) -> EventBacktestResponse:
"""自然语言事件回测:入场条件命中 -> 次日买入 -> 持有 N 日,单股或全市场汇总统计。
直传 spec 则跳过 LLM前端调参重跑"""
try:
spec = req.spec or await parse_event_spec(req.text)
result = await run_event_backtest(
session, spec,
ts_code=req.ts_code,
start=req.start.date() if req.start else None,
end=req.end.date() if req.end else None,
)
except ScreenerError as e:
raise HTTPException(status_code=502, detail=str(e))
except EventEngineError as e:
raise HTTPException(status_code=400, detail=str(e))
except Exception as e: # noqa: BLE001
raise HTTPException(status_code=500, detail=f"事件回测失败: {e}")
return EventBacktestResponse(
text=req.text,
spec=result["spec"],
universe=result["universe"],
start=result["start"],
end=result["end"],
stats=result["stats"],
trades=result["trades"],
total=result["total"],
)
# ---------- 智能选股 ----------
@router.post("/screener/run", response_model=ScreenerRunResponse)
async def screener_run(
req: ScreenerRunRequest, session: AsyncSession = Depends(get_session)
req: ScreenerRunRequest,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> ScreenerRunResponse:
"""自然语言 -> LLM 解析条件 -> 全市场筛选。也可直传 conditions 跳过 LLM微调再跑"""
"""自然语言 -> LLM 解析条件 -> 全市场筛选。也可直传 conditions 跳过 LLM微调再跑
成功的提问(含解析出的条件与命中数)记录到 screener_queries供历史一键重跑。"""
try:
conds = req.conditions or await parse_conditions(req.text)
if not conds.indicator and not conds.snapshot:
raise HTTPException(status_code=400, detail="AI 未从描述中解析出任何筛选条件,请换种说法")
result = await engine.run_screen(session, conds, settings.screener_default_limit)
# 相同文本 + 相同条件的上一条不重复记录(一键重跑场景)
exists = (
await session.execute(
select(ScreenerQuery.id).where(
ScreenerQuery.user_id == user.id,
ScreenerQuery.text == req.text.strip(),
ScreenerQuery.conditions_json == json.dumps(conds.model_dump(), ensure_ascii=False),
)
)
).scalar_one_or_none()
if exists is None:
session.add(ScreenerQuery(
user_id=user.id,
text=req.text.strip(),
conditions_json=json.dumps(conds.model_dump(), ensure_ascii=False),
hit_count=result.get("total", 0),
))
await session.commit()
return ScreenerRunResponse(**result)
except HTTPException:
raise
@@ -202,6 +411,148 @@ async def screener_run(
raise HTTPException(status_code=code, detail=str(e))
@router.get("/screener/queries", response_model=ScreenerQueryListResponse)
async def screener_queries(
limit: int = 20,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> ScreenerQueryListResponse:
"""当前用户的提问历史(最新在前,含解析出的条件与命中数,可一键重跑)。"""
limit = max(1, min(limit, 100))
rows = (
await session.execute(
select(ScreenerQuery)
.where(ScreenerQuery.user_id == user.id)
.order_by(ScreenerQuery.created_at.desc())
.limit(limit)
)
).scalars().all()
items = []
for r in rows:
conds = None
if r.conditions_json:
try:
from .schemas import ScreenConditions
conds = ScreenConditions.model_validate_json(r.conditions_json)
except Exception: # noqa: BLE001 —— 旧格式/解析失败则只展示文本
conds = None
items.append(ScreenerQueryOut(
id=r.id, text=r.text, conditions=conds, hit_count=r.hit_count, created_at=r.created_at
))
return ScreenerQueryListResponse(items=items)
@router.delete("/screener/queries/{query_id}", status_code=204)
async def screener_query_delete(
query_id: int,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> None:
await session.execute(
text("DELETE FROM screener_queries WHERE id = :i AND user_id = :u"),
{"i": query_id, "u": user.id},
)
await session.commit()
# ---------- 用户偏好 ----------
@router.get("/preferences", response_model=PreferencesOut)
async def get_preferences(
session: AsyncSession = Depends(get_session), user=Depends(require_user)
) -> PreferencesOut:
prefs: dict[str, object] = {}
rows = (
await session.execute(select(UserPreference).where(UserPreference.user_id == user.id))
).scalars().all()
for r in rows:
try:
prefs[r.key] = json.loads(r.value_json)
except Exception: # noqa: BLE001
prefs[r.key] = None
return PreferencesOut(prefs=prefs)
@router.put("/preferences", response_model=PreferencesOut)
async def put_preferences(
req: PreferencesUpdate,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> PreferencesOut:
"""部分更新:只覆盖出现的 key值为 null 表示删除该 key。返回更新后的全量。"""
for key, value in req.prefs.items():
if not key or len(key) > 64:
continue
if value is None:
await session.execute(
text("DELETE FROM user_preferences WHERE user_id = :u AND key = :k"),
{"u": user.id, "k": key},
)
continue
existing = (
await session.execute(
select(UserPreference).where(
UserPreference.user_id == user.id, UserPreference.key == key
)
)
).scalars().first()
vj = json.dumps(value, ensure_ascii=False)
if existing:
existing.value_json = vj
else:
session.add(UserPreference(user_id=user.id, key=key, value_json=vj))
await session.commit()
return await get_preferences(session=session, user=user)
# ---------- 自选股 ----------
@router.get("/watchlist", response_model=list[str])
async def get_watchlist(
session: AsyncSession = Depends(get_session), user=Depends(require_user)
) -> list[str]:
"""当前用户自选股 ts_code 列表(加入时间倒序)。"""
rows = (
await session.execute(
select(WatchlistItem.ts_code)
.where(WatchlistItem.user_id == user.id)
.order_by(WatchlistItem.created_at.desc(), WatchlistItem.id.desc())
)
).scalars().all()
return list(rows)
@router.post("/watchlist", response_model=list[str])
async def add_watchlist(
req: WatchlistOp,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> list[str]:
exists = (
await session.execute(
select(WatchlistItem.id).where(
WatchlistItem.user_id == user.id, WatchlistItem.ts_code == req.ts_code
)
)
).scalar_one_or_none()
if exists is None:
session.add(WatchlistItem(user_id=user.id, ts_code=req.ts_code))
await session.commit()
return await get_watchlist(session=session, user=user)
@router.delete("/watchlist/{ts_code}", response_model=list[str])
async def remove_watchlist(
ts_code: str,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> list[str]:
await session.execute(
text("DELETE FROM watchlist_items WHERE user_id = :u AND ts_code = :c"),
{"u": user.id, "c": ts_code},
)
await session.commit()
return await get_watchlist(session=session, user=user)
@router.post("/screener/sync", response_model=ScreenerSyncStatus)
async def screener_sync_start(
req: ScreenerSyncRequest, session: AsyncSession = Depends(get_session)
@@ -224,10 +575,22 @@ async def screener_sync_status(session: AsyncSession = Depends(get_session)) ->
@router.get("/screener/preview/{ts_code}", response_model=PreviewResponse)
async def screener_preview(
ts_code: str, limit: int = 260, session: AsyncSession = Depends(get_session)
ts_code: str, limit: int = 500, adjust: str = "qfq", timeframe: str = "1d", mas: str = "5,10,20,60",
session: AsyncSession = Depends(get_session),
) -> PreviewResponse:
"""个股详情预览:日线(qfq 全量缓存,未缓存/过期自动拉取,失败退 market_daily 近段)
+ 全套指标indicators.py 单一事实源)+ 最新截面信息卡。"""
"""个股详情预览:日线(candles 不复权底座 + adj_factor 本地换算 bfq/qfq/hfq
未缓存自动拉取,失败退 market_daily 近段)+ 全套指标 + 最新截面信息卡。
timeframe 聚合到周/月/年先复权再聚合mas 指定主图 MA 周期(逗号分隔)。"""
if adjust not in _ADJUST_MODES:
raise HTTPException(status_code=400, detail=f"adjust 仅支持 {'/'.join(_ADJUST_MODES)}")
if timeframe not in ("1d", "1w", "1M", "1y"):
raise HTTPException(status_code=400, detail="timeframe 仅支持 1d/1w/1M/1y")
try:
ma_periods = sorted({int(p) for p in mas.split(",") if p.strip().isdigit() and 1 <= int(p) <= 500})
except ValueError:
raise HTTPException(status_code=400, detail="mas 格式应为逗号分隔的数字,如 5,10,20,60")
if not ma_periods:
ma_periods = [5, 10, 20, 60]
symbol = plain_code(ts_code)
# 先取 market_daily 最新行:既做缓存过期判断,也做信息卡数据源
@@ -237,16 +600,20 @@ async def screener_preview(
)
).scalars().first()
# --- 日线candles(qfq 全量) 优先;未缓存拉取,缓存落后于全市场最新交易日则强制刷新(每日至多一次) ---
# --- 日线candles(不复权底座) 优先;未缓存拉取,缓存落后于全市场最新交易日则强制刷新(每日至多一次) ---
# fetcher 增量拉取写入的是 qfqsettings.data_adjust此时底座模式记为 qfq。
rows = await repository.get_candles(session, symbol, "1d", limit=100000)
source = "qfq"
source = "bfq"
mode = "bfq"
try:
if not rows:
await fetcher.sync_symbol(session, symbol, source="auto")
rows = await repository.get_candles(session, symbol, "1d", limit=100000)
mode = settings.data_adjust if settings.data_adjust in _ADJUST_MODES else "qfq"
elif md is not None and rows and rows[-1].ts.date() < md.trade_date.date():
await fetcher.sync_symbol(session, symbol, source="auto", force=True)
rows = await repository.get_candles(session, symbol, "1d", limit=100000)
mode = settings.data_adjust if settings.data_adjust in _ADJUST_MODES else "qfq"
except Exception: # noqa: BLE001 —— tushare/写库失败时回滚会话(否则毒化后兜底查询 500
await session.rollback()
if not rows:
@@ -265,6 +632,22 @@ async def screener_preview(
if not bars:
raise HTTPException(status_code=404, detail=f"无数据: {ts_code}(可先点「同步市场数据」)")
# --- 复权换算:请求模式与底座模式不同时按 adj_factor 本地换算(无因子则维持原样) ---
if adjust != mode:
factors = (
await session.execute(
select(AdjFactor).where(AdjFactor.ts_code == ts_code).order_by(AdjFactor.trade_date)
)
).scalars().all()
if factors:
bars = _adjust_bars(bars, factors, mode, adjust)
mode = adjust
if source != "market":
source = adjust
# --- 周期聚合:复权之后按日历聚合到周/月/年,指标在聚合后的序列上计算 ---
bars = resample_bars(bars, timeframe)
# --- 指标(在全量历史上计算后截尾,保证预热正确) ---
df = pd.DataFrame({"close": [b.close for b in bars], "high": [b.high for b in bars], "low": [b.low for b in bars]})
closes, highs, lows = df["close"], df["high"], df["low"]
@@ -272,7 +655,7 @@ async def screener_preview(
kdj = ind.kdj(highs, lows, closes)
boll = ind.bollinger(closes)
indicators: dict[str, dict[str, list[float | None]]] = {
"ma": {f"ma{p}": _series_to_jsonable(ind.ma(closes, p)) for p in (5, 10, 20, 60)},
"ma": {f"ma{p}": _series_to_jsonable(ind.ma(closes, p)) for p in ma_periods},
"macd": {
"dif": _series_to_jsonable(macd["macd"]),
"dea": _series_to_jsonable(macd["signal"]),

View File

@@ -1,6 +1,6 @@
"""数据编排拉取Tushare 主 -> AKShare 兜底)+ 本地缓存。
真实行情落库到 candles 表timeframe='1d'),回测统一从库读,与 DEMO 同路径
真实行情落库到 candles 表timeframe='1d'),回测统一从库读。
"""
from __future__ import annotations

View File

@@ -1,78 +0,0 @@
"""合成数据MVP 零依赖可跑)。
生成随机游走 OHLCV灌入 DB。仅用于让回测链路在没有真实数据源时也能跑通演示。
阶段1 接 Tushare/AKShare 后,这里仅保留为"离线测试夹具"
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
import numpy as np
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from ..domain import Bar
from ..models import Candle
def _trading_days(n: int) -> list[datetime]:
"""粗略生成 n 个工作日跳过周末节假日由阶段1 的交易日历服务处理)。"""
start = datetime.now(timezone.utc).replace(tzinfo=None) - timedelta(days=int(n * 1.6))
days: list[datetime] = []
d = start
while len(days) < n:
if d.weekday() < 5:
days.append(d.replace(hour=15, minute=0, second=0, microsecond=0))
d += timedelta(days=1)
return days
def generate_ohlcv(n: int = 500, seed: int = 42) -> list[Bar]:
"""随机游走 + A 股风格的价格区间5~30 元)。"""
rng = np.random.default_rng(seed)
rets = rng.normal(loc=0.0003, scale=0.018, size=n)
price = 10.0 * np.cumprod(1 + rets)
days = _trading_days(n)
bars: list[Bar] = []
for i in range(n):
close = float(price[i])
op = close * (1 + rng.normal(0, 0.005))
hi = max(op, close) * (1 + abs(rng.normal(0, 0.006)))
lo = min(op, close) * (1 - abs(rng.normal(0, 0.006)))
vol = float(rng.integers(1_000_000, 10_000_000))
bars.append(
Bar(
ts=days[i],
open=round(op, 2),
high=round(hi, 2),
low=round(lo, 2),
close=round(close, 2),
volume=vol,
)
)
return bars
async def seed_if_empty(session: AsyncSession, symbol: str = "DEMO", n: int = 500) -> None:
"""若库中无该 symbol 数据,则灌入合成数据。"""
existing = await session.execute(
select(Candle.id).where(Candle.symbol == symbol).limit(1)
)
if existing.scalars().first() is not None:
return
bars = generate_ohlcv(n=n)
for b in bars:
session.add(
Candle(
symbol=symbol,
timeframe="1d",
ts=b.ts,
open=b.open,
high=b.high,
low=b.low,
close=b.close,
volume=b.volume,
)
)
await session.commit()

View File

@@ -17,7 +17,7 @@ def _parse(date_str: str) -> datetime:
def fetch_daily(code: str, start: str = "20200101", end: str | None = None,
adjust: str = "qfq") -> list[Bar]:
import tushare as ts # 延迟导入:未装/无 token 时 DEMO 仍可用
import tushare as ts # 延迟导入:未装无 token 时该数据源不可用
if not settings.tushare_token:
raise RuntimeError("未配置 TUSHARE_TOKEN")

View File

@@ -21,7 +21,7 @@ async def lifespan(app: FastAPI):
app = FastAPI(
title="Stock Backtest",
description="历史回测 + 回放式模拟平台A 股为主,不做实盘)",
description="股票研究平台:全市场数据 + 智能选股 + 事件回测A 股为主,不做实盘)",
version="0.1.0",
lifespan=lifespan,
docs_url="/docs" if settings.expose_api_docs else None,

View File

@@ -9,7 +9,7 @@ Candle 表设计与 TimescaleDB hypertable 完全兼容:将来在目标 PG 库
"""
from datetime import datetime
from sqlalchemy import BigInteger, Boolean, DateTime, Float, ForeignKey, Integer, String, UniqueConstraint
from sqlalchemy import BigInteger, Boolean, DateTime, Float, ForeignKey, Integer, String, Text, UniqueConstraint
from sqlalchemy.orm import Mapped, mapped_column, relationship
from .db import Base
@@ -124,6 +124,65 @@ class DailySnapshot(Base):
)
class AdjFactor(Base):
"""复权因子adj_factorTushare 原始值qfq/hfq 本地换算的底座)。
与 candles(不复权日线) 按 ts_code+trade_date 关联:
前复权 qfq = 不复权价 × f(t) / f(latest);后复权 hfq = 不复权价 × f(t)。
"""
__tablename__ = "adj_factor"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
trade_date: Mapped[datetime] = mapped_column(DateTime, index=True)
ts_code: Mapped[str] = mapped_column(String(12), index=True)
adj_factor: Mapped[float] = mapped_column(Float)
__table_args__ = (
UniqueConstraint("ts_code", "trade_date", name="uq_adj_code_date"),
)
class UserPreference(Base):
"""用户偏好键值对(配色/复权口径/MA 周期/副图布局等value 存 JSON 字符串)。"""
__tablename__ = "user_preferences"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
user_id: Mapped[int] = mapped_column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), index=True)
key: Mapped[str] = mapped_column(String(64))
value_json: Mapped[str] = mapped_column(Text, default="null")
updated_at: Mapped[datetime] = mapped_column(DateTime, default=_utcnow, onupdate=_utcnow)
__table_args__ = (
UniqueConstraint("user_id", "key", name="uq_user_pref_key"),
)
class WatchlistItem(Base):
"""自选股(星标置顶)。"""
__tablename__ = "watchlist_items"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
user_id: Mapped[int] = mapped_column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), index=True)
ts_code: Mapped[str] = mapped_column(String(12), index=True)
created_at: Mapped[datetime] = mapped_column(DateTime, default=_utcnow)
__table_args__ = (
UniqueConstraint("user_id", "ts_code", name="uq_watch_user_code"),
)
class ScreenerQuery(Base):
"""自然语言选股提问历史(文本 + 解析出的条件,便于一键重跑)。"""
__tablename__ = "screener_queries"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
user_id: Mapped[int] = mapped_column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), index=True)
text: Mapped[str] = mapped_column(String(500))
conditions_json: Mapped[str | None] = mapped_column(Text)
hit_count: Mapped[int | None] = mapped_column(Integer)
created_at: Mapped[datetime] = mapped_column(DateTime, default=_utcnow, index=True)
class TradeCalendar(Base):
"""交易日历缓存trade_cal 拉取一次宽范围后本地维护,低积分 token 限频 1 次/小时)。"""
__tablename__ = "trade_calendar"

View File

@@ -24,7 +24,7 @@ class CandleOut(BaseModel):
# ---------- Backtest ----------
class BacktestRequest(BaseModel):
symbol: str = "DEMO"
symbol: str = "000001"
timeframe: str = "1d"
strategy: str = "macd_cross" # macd_cross | ma_cross | single_ma
params: dict[str, float] = Field(default_factory=dict) # 各策略参数
@@ -82,6 +82,68 @@ class SyncRequest(BaseModel):
force: bool = False # True => 忽略缓存重新拉取
# ---------- Event Backtest自然语言事件回测 ----------
class EventBacktestSpec(BaseModel):
"""事件回测参数entry 条件在信号日 D 收盘确认 -> D+1 买入 -> 持有 N 日卖出。"""
entry: ScreenConditions
entry_timing: Literal["next_open", "next_close"] = "next_open" # 次日开盘/收盘买入
holding_days: int = Field(default=3, ge=1, le=250) # 买入后再持有 N 个交易日
exit_timing: Literal["close", "open"] = "close" # 到期按收盘/开盘卖出
class EventBacktestRequest(BaseModel):
text: str = Field(min_length=2, max_length=500)
spec: EventBacktestSpec | None = None # 直传则跳过 LLM 解析(调参重跑)
ts_code: str | None = None # 指定则只回测该股;空则全市场
start: datetime | None = None
end: datetime | None = None
class EventTradeOut(BaseModel):
ts_code: str
name: str | None = None
entry_date: datetime
entry_price: float
exit_date: datetime
exit_price: float
ret_pct: float # 区间收益率 %(复权校正)
class EventYearStatOut(BaseModel):
year: int
samples: int
mean_pct: float
median_pct: float
win_rate: float
class EventStatsOut(BaseModel):
samples: int
stocks: int
mean_pct: float
median_pct: float
win_rate: float # %
std_pct: float = 0.0
p10_pct: float = 0.0
p25_pct: float = 0.0
p75_pct: float = 0.0
p90_pct: float = 0.0
max_pct: float = 0.0
min_pct: float = 0.0
by_year: list[EventYearStatOut] = Field(default_factory=list)
class EventBacktestResponse(BaseModel):
text: str
spec: EventBacktestSpec
universe: str # "all" 或 ts_code
start: datetime
end: datetime
stats: EventStatsOut
trades: list[EventTradeOut] = Field(default_factory=list) # 最好+最差样本(各 100
total: int
class SyncResponse(BaseModel):
symbol: str
bars: int
@@ -197,7 +259,7 @@ class PreviewInfoOut(BaseModel):
class PreviewResponse(BaseModel):
ts_code: str
symbol: str
source: str # qfq=回测缓存全量前复权 | market=近段未复权兜底
source: str # bfq|qfq|hfq=实际复权口径(本地 adj_factor 换算) | market=近段未复权兜底
info: PreviewInfoOut
candles: list[CandleOut]
indicators: dict[str, dict[str, list[float | None]]] = Field(default_factory=dict)
@@ -220,3 +282,61 @@ class CurrentUserOut(BaseModel):
class LoginResponse(BaseModel):
user: CurrentUserOut
expires_at: datetime
# ---------- 股票列表(全市场浏览) ----------
class StockListItemOut(BaseModel):
ts_code: str
symbol: str
name: str
industry: str | None = None
market: str | None = None
close: float | None = None # 最新收盘candles 未复权)
prev_close: float | None = None
pct_chg: float | None = None # 最新两根日线计算
last_ts: datetime | None = None
bar_count: int | None = None # 本地缓存日线条数
watched: bool = False # 是否自选(当前用户)
class StockListResponse(BaseModel):
total: int
items: list[StockListItemOut]
# ---------- 看股页筛选项 ----------
class FacetItemOut(BaseModel):
name: str
count: int
class StockFacetsResponse(BaseModel):
industries: list[FacetItemOut] = Field(default_factory=list)
areas: list[FacetItemOut] = Field(default_factory=list)
# ---------- 用户偏好 / 自选股 / 提问历史 ----------
class PreferencesOut(BaseModel):
prefs: dict[str, object] = Field(default_factory=dict) # key -> JSON 值
class PreferencesUpdate(BaseModel):
prefs: dict[str, object] # 部分更新:只覆盖出现的 key值为 null 表示删除)
class WatchlistOp(BaseModel):
ts_code: str = Field(min_length=6, max_length=12)
class ScreenerQueryOut(BaseModel):
id: int
text: str
conditions: ScreenConditions | None = None
hit_count: int | None = None
created_at: datetime
model_config = {"from_attributes": True}
class ScreenerQueryListResponse(BaseModel):
items: list[ScreenerQueryOut]

View File

@@ -13,7 +13,7 @@ import re
import httpx
from ..config import settings
from ..schemas import ScreenConditions
from ..schemas import EventBacktestSpec, ScreenConditions
SYSTEM_PROMPT = """你是 A 股选股条件解析器。把用户的自然语言解析成一个 JSON 对象,只输出 JSON不要任何解释、注释或代码块围栏。完全无法理解时输出 {"error": "原因"}。
@@ -149,3 +149,69 @@ async def parse_conditions(text: str) -> ScreenConditions:
except Exception as e: # noqa: BLE001 —— JSON/校验失败,带错误重试
retry_error = str(e)[:300]
raise ScreenerError(f"AI 解析结果两次未通过校验,最后错误:{retry_error}")
# ---------- 事件回测解析(自然语言 -> EventBacktestSpec ----------
EVENT_SYSTEM_PROMPT = """你是 A 股事件回测参数解析器。用户描述一个「入场信号 + 买卖时机 + 持有期」的事件回测需求,把它解析成 JSON只输出 JSON不要任何解释或代码块围栏。完全无法理解时输出 {"error": "原因"}。
输出结构:
{"entry": {"indicator": [...], "snapshot": [], "exclude_st": true, "exclude_delisted": true, "exclude_bj": true}, "entry_timing": "next_open", "holding_days": 3, "exit_timing": "close"}
【entry.indicator 数组】入场信号条件(必填,至少 1 条),元素字段与白名单:
- "indicator": kdj_k / kdj_d / kdj_jKDJ 的 K/D/J 值、rsi、macd_dif / macd_dea / macd_histMACD 的 DIF/DEA/柱、ma收盘价均线、boll_upper / boll_mid / boll_lower布林轨道、close收盘价、pct_chg日涨跌幅%
- "params": 指标参数可选默认KDJ {"n":9,"m1":3,"m2":3}RSI {"period":14}MACD {"fast":12,"slow":26,"signal":9}MA {"period":20}BOLL {"period":20,"std":2}
- "op": "gt" | "ge" | "lt" | "le" | "between""value"between 时为下界)、"value2"(上界)
- "value_indicator": 指标与指标比较时填另一指标名(同白名单),如 "DIF 大于 DEA" -> indicator=macd_dif, op=gt, value_indicator=macd_dea, value=0
- "value_params": 比较对象指标参数不同时指定,如 "MA5 上穿 MA20" -> indicator=ma, params={"period":5}, op=gt, value_indicator=ma, value_params={"period":20}, value=0
- "lookback": 信号需连续/曾经满足的交易日窗口(默认 1
- "match": "all"(窗口内每天满足,默认)或 "any"(窗口内任一天满足)
【时间语义】"连续三天 J 小于 10" -> lookback=3, match="all""近 5 天曾经金叉" -> lookback=5, match="any"
【entry_timing】买入时机"第二天开盘购买/次日开盘买入" -> "next_open"(默认);"第二天收盘买入" -> "next_close"
【holding_days】买入后持有 N 个交易日int默认 3"未来三天的涨幅" -> holding_days=3"持有 10 天" -> 10"持有一个月" -> 20。
【exit_timing】到期卖出价"close"(收盘卖,默认)或 "open"(开盘卖)。
【entry.snapshot】截面过滤条件一般不适用于历史回测除非用户明确说"只回测市值大于 X 亿的股票"才填,其余情况留空数组。
示例:
输入:在连续三天 J 小于 10 的时候第二天开盘购买,之后未来三天的涨幅有多少
输出:{"entry":{"indicator":[{"indicator":"kdj_j","params":{"n":9,"m1":3,"m2":3},"op":"lt","value":10,"lookback":3,"match":"all"}],"snapshot":[],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true},"entry_timing":"next_open","holding_days":3,"exit_timing":"close"}
示例:
输入RSI 低于 30 的第二天开盘买入持有 5 天收盘卖出
输出:{"entry":{"indicator":[{"indicator":"rsi","params":{"period":14},"op":"lt","value":30,"lookback":1,"match":"all"}],"snapshot":[],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true},"entry_timing":"next_open","holding_days":5,"exit_timing":"close"}"""
def _build_event_messages(text: str, retry_error: str | None = None) -> list[dict]:
user = f"解析以下事件回测需求:{text}"
if retry_error:
user += f"\n\n上一次输出无法通过校验,错误:{retry_error}。请修正后重新只输出 JSON。"
return [{"role": "system", "content": EVENT_SYSTEM_PROMPT}, {"role": "user", "content": user}]
async def parse_event_spec(text: str) -> EventBacktestSpec:
"""自然语言 -> EventBacktestSpec。复用 _chat/_extract_json失败带错误重试 1 次。"""
if not settings.llm_api_key:
raise ScreenerError(
"未配置 LLM_API_KEY请在 backend/.env 填入 DeepSeek API Keyplatform.deepseek.com 获取)后重启后端"
)
retry_error: str | None = None
for _ in range(2):
content = await _chat(_build_event_messages(text, retry_error))
try:
obj = _extract_json(content)
if "error" in obj and not obj.get("entry"):
raise ScreenerError(f"AI 无法理解该回测需求:{obj['error']}")
spec = EventBacktestSpec.model_validate(obj)
if not spec.entry.indicator:
raise ValueError("entry.indicator 不能为空")
return spec
except ScreenerError:
raise
except Exception as e: # noqa: BLE001
retry_error = str(e)[:300]
raise ScreenerError(f"AI 解析回测参数两次未通过校验,最后错误:{retry_error}")

View File

@@ -16,7 +16,7 @@ from sqlalchemy import delete, func, insert, select
from sqlalchemy.ext.asyncio import AsyncSession
from ..config import settings
from ..models import DailySnapshot, MarketDaily, StockBasic, TradeCalendar
from ..models import AdjFactor, DailySnapshot, MarketDaily, StockBasic, TradeCalendar
from .llm import ScreenerError
# 进程内单例任务状态uvicorn --reload 单进程场景够用)
@@ -160,6 +160,18 @@ def _fetch_basic(pro, d: str) -> list[dict]:
return rows
def _fetch_adj_factor(pro, d: str) -> list[dict]:
"""拉取某交易日全市场复权因子K线 bfq->qfq/hfq 本地换算的底座)。"""
time.sleep(settings.screener_sync_interval)
df = _call_retry(pro.adj_factor, trade_date=d)
if df is None or df.empty:
return []
return [
{"trade_date": _parse_d(d), "ts_code": r["ts_code"], "adj_factor": float(r["adj_factor"])}
for _, r in df.iterrows()
]
def _sync_stock_list_sync(pro) -> list[dict]:
"""拉取在市股票列表。"""
time.sleep(settings.screener_sync_interval)
@@ -243,6 +255,16 @@ async def _run_sync(days: int, force: bool) -> None:
await _replace_day(session, MarketDaily, daily_rows, d)
_sync_state["done_days"] += 1
# 2.5) 复权因子(与日线同窗口增量;历史全量由 scripts/backfill_adj_factor.py 回补)
async with async_session() as session:
have_adj = set() if force else await _existing_dates(session, AdjFactor)
for d in [d for d in dates if d not in have_adj]:
_sync_state["step"] = f"正在同步 {d} 复权因子"
adj_rows = await asyncio.to_thread(_fetch_adj_factor, pro, d)
if adj_rows:
async with async_session() as session:
await _replace_day(session, AdjFactor, adj_rows, d)
# 3) 最新「有数据」交易日的快照daily_basic仅 1 次调用)
# 用 market_daily 实际最大交易日(今天的数据收盘后才生成,日历最新日会拉到空)
async with async_session() as session: