功能更新
This commit is contained in:
@@ -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_mode(bfq/qfq/hfq)。
|
||||
|
||||
相对不复权的乘数:bfq=1,qfq=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 增量拉取写入的是 qfq(settings.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"]),
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""数据编排:拉取(Tushare 主 -> AKShare 兜底)+ 本地缓存。
|
||||
|
||||
真实行情落库到 candles 表(timeframe='1d'),回测统一从库读,与 DEMO 同路径。
|
||||
真实行情落库到 candles 表(timeframe='1d'),回测统一从库读。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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_factor,Tushare 原始值;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"
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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_j(KDJ 的 K/D/J 值)、rsi、macd_dif / macd_dea / macd_hist(MACD 的 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 Key(platform.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}")
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user