看股功能更新

This commit is contained in:
2026-08-16 00:05:26 +08:00
parent 9cce670b74
commit fc86fe0674
28 changed files with 3823 additions and 96 deletions

View File

@@ -15,11 +15,13 @@ import json
from datetime import datetime
import pandas as pd
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy import select, text
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile
from sqlalchemy import delete, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql.elements import TextClause
from .backtest.engine import BacktestConfig, run_backtest
from . import cache
from .auth import require_user
from .backtest.events import EventEngineError, run_event_backtest
from .backtest.strategies import build_strategy
@@ -30,6 +32,7 @@ from .data.symbols import plain_code
from .db import get_session
from .domain import Bar
from . import indicators as ind
from .trades import parse_statement
from .models import (
AdjFactor,
BacktestRun,
@@ -38,6 +41,7 @@ from .models import (
ScreenerQuery,
StockBasic,
UserPreference,
UserTrade,
WatchlistItem,
)
from .schemas import (
@@ -66,6 +70,9 @@ from .schemas import (
FacetItemOut,
SyncRequest,
SyncResponse,
TradesClearResponse,
TradesImportResponse,
UserTradeOut,
WatchlistOp,
)
from .screener import engine, market_sync
@@ -163,37 +170,63 @@ async def sync_data(req: SyncRequest, session: AsyncSession = Depends(get_sessio
# ---------- 股票列表(全市场浏览) ----------
_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
# 过滤/排序/分页在 stock_basic+watchlist+daily_snapshot 上完成(快照按最新交易日
# 走唯一索引 join便宜再对「本页」≤limit 只股票补最新价/昨收LATERAL 扫
# candles——旧写法对全市场 ~5000 只逐个算,每页都白算 50 倍的行情量。
# 排序列白名单键→CTE 内表达式order_by 由白名单拼接进模板,不接收用户原文。
_STOCKS_SORTS = {
"symbol": "sb.symbol",
"total_mv": "snap.total_mv",
"circ_mv": "snap.circ_mv",
"pe_ttm": "snap.pe_ttm",
"pb": "snap.pb",
"turnover_rate": "snap.turnover_rate",
}
_STOCKS_SQL_TMPL = """
WITH page AS (
SELECT sb.ts_code, sb.symbol, sb.name, sb.industry, sb.market,
(w.id IS NOT NULL) AS watched,
snap.turnover_rate, snap.pe_ttm, snap.pb, snap.total_mv, snap.circ_mv
FROM stock_basic sb
LEFT JOIN watchlist_items w ON w.ts_code = sb.ts_code AND w.user_id = :uid
LEFT JOIN daily_snapshot snap ON snap.ts_code = sb.ts_code
AND snap.trade_date = (SELECT max(trade_date) FROM daily_snapshot)
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 {order_by}
LIMIT :limit OFFSET :offset
)
SELECT p.ts_code, p.symbol, p.name, p.industry, p.market, p.watched,
c.close AS close, prev.close AS prev_close, c.ts AS last_ts,
CASE WHEN c.close IS NOT NULL AND prev.close IS NOT NULL AND prev.close <> 0
THEN round(((c.close / prev.close - 1) * 100)::numeric, 2) END AS pct_chg,
p.turnover_rate, p.pe_ttm, p.pb,
round((p.total_mv / 10000.0)::numeric, 2) AS total_mv,
round((p.circ_mv / 10000.0)::numeric, 2) AS circ_mv
FROM page p
LEFT JOIN LATERAL (
SELECT close, ts FROM candles
WHERE symbol = sb.symbol AND timeframe = '1d'
WHERE symbol = p.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
WHERE symbol = p.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
""")
) prev ON c.ts IS NOT NULL
"""
def _stocks_sql(sort: str, order: str) -> TextClause:
col = _STOCKS_SORTS.get(sort, _STOCKS_SORTS["symbol"])
direction = "DESC" if order == "desc" else "ASC"
nulls = " NULLS LAST" if col != "sb.symbol" else "" # 快照缺失/亏损无 PE 的排最后
return text(_STOCKS_SQL_TMPL.format(order_by=f"{col} {direction}{nulls}"))
_STOCKS_COUNT_SQL = text("""
SELECT count(*) FROM stock_basic sb
@@ -214,16 +247,32 @@ async def list_stocks(
industry: str = "",
area: str = "",
watched_only: bool = False,
sort: str = "symbol",
order: str = "asc",
limit: int = 100,
offset: int = 0,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> StockListResponse:
"""全市场股票列表stock_basic 基本信息 + candles 最新行情(本地缓存,无缓存则行情列为空)。
自选股watchlist_items排最前watched_only=true 只看自选。"""
"""全市场股票列表stock_basic 基本信息 + candles 最新行情 + daily_snapshot 估值指标
(换手率/PE-TTM/PB/市值,无快照则这些列为空)。
watched_only=true 只看自选(自选有独立的「自选」分类入口,列表不再把自选排最前)。
sort ∈ {symbol,total_mv,circ_mv,pe_ttm,pb,turnover_rate}(白名单,其他值回落 symbol
order ∈ asc/desc快照列排序时缺失值无快照/亏损无 PE恒排末尾。
Redis 缓存:按「用户自选版本 + 查询参数(含排序)」缓存整页(含 total自选增删即时失效。"""
search = search.strip()
sort = sort if sort in _STOCKS_SORTS else "symbol"
order = "desc" if order.lower() == "desc" else "asc"
limit = max(1, min(limit, 500))
offset = max(0, offset)
key = (
f"stocks:u{user.id}"
f":v{await cache.get_version(f'watchlist:{user.id}')}"
f":{cache.digest(search, market, industry, area, watched_only, sort, order, limit, offset)}"
)
cached = await cache.cache_get(key)
if cached is not None:
return StockListResponse(**cached)
params = {
"search": search,
"psearch": f"%{search}%",
@@ -236,13 +285,18 @@ async def list_stocks(
"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])
rows = (await session.execute(_stocks_sql(sort, order), params)).mappings().all()
resp = StockListResponse(total=total, items=[StockListItemOut(**r) for r in rows])
await cache.cache_set(key, resp.model_dump(mode="json"), settings.stocks_cache_ttl)
return resp
@router.get("/stocks/facets", response_model=StockFacetsResponse)
async def stock_facets(session: AsyncSession = Depends(get_session)) -> StockFacetsResponse:
"""看股页筛选项:行业 / 地域(含数量,按数量降序)。"""
"""看股页筛选项:行业 / 地域(含数量,按数量降序)。stock_basic 很少变,长缓存。"""
cached = await cache.cache_get("facets:stocks")
if cached is not None:
return StockFacetsResponse(**cached)
industries = (
await session.execute(text("""
SELECT industry AS name, count(*) AS n FROM stock_basic
@@ -257,10 +311,12 @@ async def stock_facets(session: AsyncSession = Depends(get_session)) -> StockFac
GROUP BY area ORDER BY n DESC
"""))
).mappings().all()
return StockFacetsResponse(
resp = 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],
)
await cache.cache_set("facets:stocks", resp.model_dump(mode="json"), settings.facets_cache_ttl)
return resp
@router.post("/backtest", response_model=BacktestResponse)
@@ -551,6 +607,7 @@ async def add_watchlist(
if exists is None:
session.add(WatchlistItem(user_id=user.id, ts_code=req.ts_code))
await session.commit()
await cache.bump_version(f"watchlist:{user.id}") # 作废该用户的股票列表缓存
return await get_watchlist(session=session, user=user)
@@ -565,9 +622,125 @@ async def remove_watchlist(
{"u": user.id, "c": ts_code},
)
await session.commit()
await cache.bump_version(f"watchlist:{user.id}") # 作废该用户的股票列表缓存
return await get_watchlist(session=session, user=user)
# ---------- 交割单(个人实盘买卖点) ----------
def _trade_out(r: UserTrade) -> UserTradeOut:
return UserTradeOut(
id=r.id, ts_code=r.ts_code, name=r.name, trade_date=r.trade_date,
direction=r.direction, price=r.price, qty=r.qty, amount=r.amount, fee=r.fee,
)
@router.get("/trades", response_model=list[UserTradeOut])
async def list_trades(
ts_code: str | None = None,
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> list[UserTradeOut]:
"""当前用户导入的实盘成交(可选 ts_code 过滤按日期升序K线买卖点数据源"""
q = (
select(UserTrade)
.where(UserTrade.user_id == user.id)
.order_by(UserTrade.trade_date, UserTrade.id)
)
if ts_code:
q = q.where(UserTrade.ts_code == ts_code)
rows = (await session.execute(q)).scalars().all()
return [_trade_out(r) for r in rows]
@router.post("/trades/import", response_model=TradesImportResponse)
async def import_trades(
file: UploadFile = File(...),
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> TradesImportResponse:
"""上传券商交割单CSV/Excel/HTML 表格均可,自动识别列名),解析出买卖成交入库。
同一笔成交(同日同股同向同价同量)重复上传会跳过,重复导出幂等。
"""
data = await file.read()
if not data:
raise HTTPException(status_code=422, detail="文件是空的")
if len(data) > 20 * 1024 * 1024:
raise HTTPException(status_code=413, detail="文件超过 20MB请分时间段导出")
parsed = parse_statement(data, file.filename or "")
# 无证券代码列的导出(招商式):按证券名称反查 stock_basic 补 ts_code同名多码或查不到则弃行
unnamed = {t.name for t in parsed.trades if not t.ts_code and t.name}
if unnamed:
name_map: dict[str, str] = {}
for ts_code, name in (await session.execute(
select(StockBasic.ts_code, StockBasic.name).where(StockBasic.name.in_(unnamed))
)).all():
name_map[name] = "" if name in name_map else ts_code
for t in parsed.trades:
if not t.ts_code and t.name:
tc = name_map.get(t.name, "")
if tc:
t.ts_code, t.code = tc, tc.split(".")[0]
else:
parsed.skipped_bad.append(f"{t.trade_date} {t.name} 名称无法唯一对应代码,未入库")
def _key(t) -> tuple:
return (t.trade_date, t.ts_code, t.direction, None if t.price is None else round(t.price, 4), round(t.qty, 4))
# Python 侧去重兜底(唯一约束对 NULL price 不生效)
existing = {
(r.trade_date, r.ts_code, r.direction, None if r.price is None else round(r.price, 4), round(r.qty, 4))
for r in (
await session.execute(
select(UserTrade.trade_date, UserTrade.ts_code, UserTrade.direction, UserTrade.price, UserTrade.qty)
.where(UserTrade.user_id == user.id, UserTrade.ts_code.in_({t.ts_code for t in parsed.trades}))
)
).all()
}
inserted: list[UserTrade] = []
seen: set[tuple] = set()
skipped_dup = 0
for t in parsed.trades:
if not t.ts_code:
continue # 名称反查失败的行,已在 bad 里说明
k = _key(t)
if k in existing or k in seen:
skipped_dup += 1
continue
seen.add(k)
inserted.append(UserTrade(
user_id=user.id, ts_code=t.ts_code, code=t.code, name=t.name or None,
trade_date=t.trade_date, direction=t.direction, price=t.price,
qty=t.qty, amount=t.amount, fee=t.fee,
raw_json=json.dumps(t.raw, ensure_ascii=False, default=str),
))
if inserted:
session.add_all(inserted)
await session.commit()
return TradesImportResponse(
inserted=len(inserted),
skipped_dup=skipped_dup,
skipped_other=parsed.skipped_other,
stocks=len({t.ts_code for t in parsed.trades}),
bad=parsed.skipped_bad[:5],
sample=[_trade_out(r) for r in inserted[:5]],
)
@router.delete("/trades", response_model=TradesClearResponse)
async def clear_trades(
session: AsyncSession = Depends(get_session),
user=Depends(require_user),
) -> TradesClearResponse:
"""清空当前用户导入的全部成交(重新导入前用)。"""
res = await session.execute(delete(UserTrade).where(UserTrade.user_id == user.id))
await session.commit()
return TradesClearResponse(deleted=res.rowcount or 0)
@router.post("/screener/sync", response_model=ScreenerSyncStatus)
async def screener_sync_start(
req: ScreenerSyncRequest, session: AsyncSession = Depends(get_session)

View File

@@ -0,0 +1,265 @@
"""事件回测引擎:入场条件命中 -> 次日买入 -> 持有 N 日 -> 全市场汇总统计。
数据口径:
- 行情底座是 candlesTDX 全量导入,不复权),全历史可用;
- 指标计算用不复权价(与选股/看盘口径一致J<10、RSI<30 等阈值均为归一化或惯例值);
- 收益率用 adj_factor 校正ret = 出场价×f出 / 入场价×f入 - 1消除除权除息失真
因子缺失的股退化为不复权收益(新股/缺因子,样本中占少数)。
信号语义:与选股引擎一致——每条条件在信号日 d 为终点、lookback 窗口内
match=all(连续满足)/any(曾经满足),多条件之间取 AND。
"""
from __future__ import annotations
from datetime import date, datetime, timedelta
import numpy as np
import pandas as pd
from sqlalchemy import and_, func, not_, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from ..models import AdjFactor, Candle, StockBasic
from ..schemas import EventBacktestSpec
from ..screener.engine import (
FAMILIES,
_family_of,
_op_mask,
_params_for,
_resolve_params,
_series_for,
)
# 指标配热缓冲 bar 数MACD 等 EMA 类指标需要较长窗口才收敛)
BUFFER_BARS = 80
# 每批查询的股票数(全市场分块拉取,避免单条 SQL 过大)
BATCH_SIZE = 800
# 单次回测允许的最大样本数(超过则仅按日期取最近的,防内存失控)
MAX_TRADES = 200_000
class EventEngineError(RuntimeError):
"""事件回测可预期的业务错误(信息透传前端)。"""
def _signal_mask(g: pd.DataFrame, spec: EventBacktestSpec, cache: dict) -> pd.Series:
"""单股全序列信号掩码各条件rolling lookbackAND。"""
total = pd.Series(True, index=g.index)
for cond in spec.entry.indicator:
fam = _family_of(cond.indicator)
if len(g) < FAMILIES[fam].min_bars:
return pd.Series(False, index=g.index)
s = _series_for(g, cond.indicator, _params_for(cond.indicator, cond.params), cache)
if s is None:
return pd.Series(False, index=g.index)
if cond.value_indicator:
target = _series_for(g, cond.value_indicator,
_resolve_params(cond, cond.value_indicator), cache)
if target is None:
return pd.Series(False, index=g.index)
else:
target = pd.Series(cond.value, index=s.index)
m = _op_mask(s, target, cond).astype(int)
n = max(1, cond.lookback)
if n > 1:
rolled = m.rolling(n, min_periods=n).sum()
m = (rolled == n) if cond.match == "all" else (rolled > 0)
else:
m = m.astype(bool)
total = total & m.fillna(False).astype(bool)
return total
def _entry_exit_indices(sig_idx: int, spec: EventBacktestSpec, n: int) -> tuple[int, int] | None:
"""信号日索引 -> (入场索引, 出场索引)。前视/越界返回 None。"""
entry_i = sig_idx + 1 # 信号收盘后才动手:一律次日
exit_i = entry_i + spec.holding_days
if exit_i >= n:
return None
return entry_i, exit_i
def _price_at(row: pd.Series, timing: str) -> float:
return float(row["open"] if timing == "open" else row["close"])
def _stats_block(trades: list[dict]) -> dict:
"""样本集合 -> 汇总统计(空样本给零值)。"""
if not trades:
return {
"samples": 0, "stocks": 0,
"mean_pct": 0.0, "median_pct": 0.0, "win_rate": 0.0, "std_pct": 0.0,
"p10_pct": 0.0, "p25_pct": 0.0, "p75_pct": 0.0, "p90_pct": 0.0,
"max_pct": 0.0, "min_pct": 0.0, "by_year": [],
}
rets = np.array([t["ret_pct"] for t in trades], dtype=float)
by_year: list[dict] = []
df = pd.DataFrame(trades)
for year, grp in df.groupby(df["entry_date"].dt.year):
r = grp["ret_pct"].to_numpy()
by_year.append({
"year": int(year), "samples": int(len(r)),
"mean_pct": round(float(r.mean()), 3),
"median_pct": round(float(np.median(r)), 3),
"win_rate": round(float((r > 0).mean() * 100), 2),
})
by_year.sort(key=lambda x: x["year"])
return {
"samples": int(len(rets)),
"stocks": int(df["ts_code"].nunique()),
"mean_pct": round(float(rets.mean()), 3),
"median_pct": round(float(np.median(rets)), 3),
"win_rate": round(float((rets > 0).mean() * 100), 2),
"std_pct": round(float(rets.std(ddof=1)) if len(rets) > 1 else 0.0, 3),
"p10_pct": round(float(np.percentile(rets, 10)), 3),
"p25_pct": round(float(np.percentile(rets, 25)), 3),
"p75_pct": round(float(np.percentile(rets, 75)), 3),
"p90_pct": round(float(np.percentile(rets, 90)), 3),
"max_pct": round(float(rets.max()), 3),
"min_pct": round(float(rets.min()), 3),
"by_year": by_year,
}
async def run_event_backtest(
session: AsyncSession,
spec: EventBacktestSpec,
ts_code: str | None = None,
start: date | None = None,
end: date | None = None,
) -> dict:
"""主入口:返回 {spec, universe, start, end, stats, trades(sample), total}。"""
entry = spec.entry
if not entry.indicator:
raise EventEngineError("入场条件必须包含技术指标条件(如 J<10、RSI<30")
# 时间窗默认最近一年end 以 candles 最大日期为准
end_dt = end
if end_dt is None:
end_dt = (await session.scalar(select(func.max(Candle.ts)))) or date.today()
if isinstance(end_dt, datetime):
end_dt = end_dt.date()
start_dt = start or (end_dt - timedelta(days=365))
if start_dt >= end_dt:
raise EventEngineError("回测起始日期必须早于结束日期")
needed = _max_needed_bars_safe(entry) + BUFFER_BARS
buffer_start = start_dt - timedelta(days=int(needed * 1.7)) # 交易日->日历日近似
# 股票池ts_code+symbol 映射candles 按 symbol 存)
name_map: dict[str, str] = {}
if ts_code:
rows = (await session.execute(
select(StockBasic.ts_code, StockBasic.symbol, StockBasic.name)
.where(StockBasic.ts_code == ts_code)
)).all()
if not rows:
raise EventEngineError(f"未知股票代码: {ts_code}")
universe = [(r[0], r[1]) for r in rows]
name_map = {r[0]: r[2] for r in rows}
else:
stmt = select(StockBasic.ts_code, StockBasic.symbol, StockBasic.name).where(
StockBasic.list_status == "L"
)
if entry.exclude_st:
stmt = stmt.where(not_(or_(StockBasic.name.like("%ST%"), StockBasic.name.like("%退%"))))
if entry.exclude_bj:
stmt = stmt.where(not_(StockBasic.ts_code.like("%.BJ")))
rows = (await session.execute(stmt)).all()
universe = [(r[0], r[1]) for r in rows]
name_map = {r[0]: r[2] for r in rows}
start_ts = datetime(start_dt.year, start_dt.month, start_dt.day)
end_ts = datetime(end_dt.year, end_dt.month, end_dt.day, 23, 59, 59)
buffer_ts = datetime(buffer_start.year, buffer_start.month, buffer_start.day)
trades: list[dict] = []
for i in range(0, len(universe), BATCH_SIZE):
batch = universe[i : i + BATCH_SIZE]
symbols = [sym for _, sym in batch]
code_by_symbol = {sym: code for code, sym in batch}
candle_rows = (await session.execute(
select(Candle.symbol, Candle.ts, Candle.open, Candle.high,
Candle.low, Candle.close)
.where(and_(Candle.timeframe == "1d",
Candle.symbol.in_(symbols),
Candle.ts >= buffer_ts, Candle.ts <= end_ts))
.order_by(Candle.symbol, Candle.ts)
)).all()
if not candle_rows:
continue
codes = {code_by_symbol[s] for s in symbols}
adj_rows = (await session.execute(
select(AdjFactor.ts_code, AdjFactor.trade_date, AdjFactor.adj_factor)
.where(and_(AdjFactor.ts_code.in_(codes),
AdjFactor.trade_date >= buffer_ts, AdjFactor.trade_date <= end_ts))
)).all()
f_map = {(r[0], r[1].date()): float(r[2]) for r in adj_rows if r[2]}
bars = pd.DataFrame(
candle_rows, columns=["symbol", "ts", "open", "high", "low", "close"]
)
for symbol, g in bars.groupby("symbol", sort=False):
if len(g) < 30:
continue
g = g.reset_index(drop=True)
ts_code_l = code_by_symbol[symbol]
cache: dict = {"_families": set()}
mask = _signal_mask(g, spec, cache)
if not mask.any():
continue
for sig_i in np.flatnonzero(mask.to_numpy()):
ts_sig = g.at[sig_i, "ts"]
# 信号必须落在回测窗口内buffer 区只用于指标配热)
if ts_sig < start_ts:
continue
ie = _entry_exit_indices(int(sig_i), spec, len(g))
if ie is None:
continue
entry_i, exit_i = ie
e_row, x_row = g.iloc[entry_i], g.iloc[exit_i]
e_price = _price_at(e_row, "open" if spec.entry_timing == "next_open" else "close")
x_price = _price_at(x_row, "open" if spec.exit_timing == "open" else "close")
if not e_price or not x_price:
continue
f_in = f_map.get((ts_code_l, e_row["ts"].date()), 1.0)
f_out = f_map.get((ts_code_l, x_row["ts"].date()), 1.0)
ret_pct = (x_price * f_out) / (e_price * f_in) * 100 - 100
trades.append({
"ts_code": ts_code_l,
"name": name_map.get(ts_code_l),
"entry_date": e_row["ts"], "entry_price": round(e_price, 3),
"exit_date": x_row["ts"], "exit_price": round(x_price, 3),
"ret_pct": round(float(ret_pct), 3),
})
if len(trades) >= MAX_TRADES:
break
if len(trades) >= MAX_TRADES:
break
if len(trades) >= MAX_TRADES:
break
stats = _stats_block(trades)
# 明细样本:最好 100 + 最差 100其余统计已覆盖
trades_sorted = sorted(trades, key=lambda t: t["ret_pct"], reverse=True)
sample = trades_sorted[:100] + (trades_sorted[-100:] if len(trades_sorted) > 100 else [])
return {
"spec": spec,
"universe": ts_code or "all",
"start": start_ts,
"end": end_ts,
"stats": stats,
"trades": sample,
"total": stats["samples"],
}
# ---------- 小工具 ----------
def _max_needed_bars_safe(conds) -> int:
"""指标配热所需最大 bar 数(同 screener.engine._max_needed_bars"""
need = 1
for c in conds.indicator:
need = max(need, FAMILIES[_family_of(c.indicator)].min_bars + c.lookback)
if c.value_indicator:
need = max(need, FAMILIES[_family_of(c.value_indicator)].min_bars + c.lookback)
return need

104
backend/app/cache.py Normal file
View File

@@ -0,0 +1,104 @@
"""Redis 读缓存(可选基础设施)。
- REDIS_URL 留空、连接失败或超时:所有操作静默退化为「无缓存」,接口照常直查数据库,
且本进程内禁用重试(避免每个请求都陪跑一次连接超时)。
- 失效策略TTL 自然过期 + 版本号INCR作废。自选股增删等写操作只 INCR 版本 key
旧缓存 key 里带着旧版本号,无需 SCAN 批量删除。
- 只缓存「读多写少、可容忍短暂陈旧」的聚合数据(股票列表、筛选项等);
K线/回测等口径敏感数据不走这里。
"""
from __future__ import annotations
import hashlib
import json
from typing import Any
import redis.asyncio as aioredis
from .config import settings
_pool: aioredis.ConnectionPool | None = None
_disabled = False # 一次失败后本进程禁用Redis 属加速件,坏了不能拖慢接口)
def _client() -> aioredis.Redis | None:
global _pool, _disabled
if not settings.redis_url or _disabled:
return None
if _pool is None:
_pool = aioredis.ConnectionPool.from_url(
settings.redis_url,
decode_responses=True,
socket_connect_timeout=1.0,
socket_timeout=1.0,
health_check_interval=60,
max_connections=32,
)
return aioredis.Redis(connection_pool=_pool)
def _bail() -> None:
global _disabled
_disabled = True
def digest(*parts: Any) -> str:
"""参数指纹(拼接后 md5仅用于拼缓存 key非安全用途"""
raw = "\x1f".join(repr(p) for p in parts)
return hashlib.md5(raw.encode()).hexdigest() # noqa: S324
async def cache_get(key: str) -> Any | None:
c = _client()
if c is None:
return None
try:
raw = await c.get(key)
return json.loads(raw) if raw is not None else None
except Exception: # noqa: BLE001 —— 缓存层任何故障都不影响主流程
_bail()
return None
async def cache_set(key: str, value: Any, ttl: int) -> None:
c = _client()
if c is None:
return
try:
await c.set(key, json.dumps(value, ensure_ascii=False), ex=max(1, ttl))
except Exception: # noqa: BLE001
_bail()
async def get_version(name: str) -> int:
"""读版本号(缺省 0。版本号参与缓存 keyINCR 后旧 key 全部失效。"""
c = _client()
if c is None:
return 0
try:
v = await c.get(f"ver:{name}")
return int(v) if v is not None else 0
except Exception: # noqa: BLE001
_bail()
return 0
async def bump_version(name: str) -> None:
c = _client()
if c is None:
return
try:
await c.incr(f"ver:{name}")
except Exception: # noqa: BLE001
_bail()
async def aclose() -> None:
"""进程退出时释放连接池(由 main.lifespan 调用)。"""
global _pool
if _pool is not None:
try:
await _pool.disconnect()
except Exception: # noqa: BLE001
pass
_pool = None

View File

@@ -26,6 +26,11 @@ class Settings(BaseSettings):
data_adjust: str = "qfq" # 复权qfq 前复权 / hfq 后复权 / "" 不复权
data_default_start: str = "20200101" # 默认拉取起点(约近 5 年)
# ---- Redis 读缓存(股票列表/筛选项等读多写少接口;留空 = 不缓存,直查数据库)----
redis_url: str = ""
stocks_cache_ttl: int = 300 # 股票列表缓存秒数(行情列允许最多滞后这么多秒)
facets_cache_ttl: int = 3600 # 行业/地域筛选项缓存秒数stock_basic 很少变)
# ---- LLM智能选股的自然语言解析DeepSeekOpenAI 兼容协议,可换任意兼容网关)----
llm_base_url: str = "https://api.deepseek.com"
llm_api_key: str = "" # 留空则智能选股不可用(其余功能不受影响)

View File

@@ -5,6 +5,7 @@ from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from sqlalchemy import text
from . import cache
from .api import router
from .auth_api import router as auth_router
from .config import settings
@@ -17,6 +18,7 @@ async def lifespan(app: FastAPI):
await conn.execute(text("SELECT 1"))
yield
await engine.dispose()
await cache.aclose() # 释放 Redis 连接池(未启用时是 no-op
app = FastAPI(

View File

@@ -7,9 +7,9 @@ Candle 表设计与 TimescaleDB hypertable 完全兼容:将来在目标 PG 库
智能选股三表stock_basic / market_daily / daily_snapshot与回测 candles(qfq)
完全隔离:选股用未复权日线按 trade_date 全市场批量落地,避免污染回测复权缓存。
"""
from datetime import datetime
from datetime import date, datetime
from sqlalchemy import BigInteger, Boolean, DateTime, Float, ForeignKey, Integer, String, Text, UniqueConstraint
from sqlalchemy import BigInteger, Boolean, Date, DateTime, Float, ForeignKey, Integer, String, Text, UniqueConstraint
from sqlalchemy.orm import Mapped, mapped_column, relationship
from .db import Base
@@ -173,6 +173,30 @@ class WatchlistItem(Base):
)
class UserTrade(Base):
"""交割单导入的实盘成交流水K线买卖点的数据源价格为券商成交原始价、不复权"""
__tablename__ = "user_trades"
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)
code: Mapped[str] = mapped_column(String(10)) # 6 位纯数字
name: Mapped[str | None] = mapped_column(String(32))
trade_date: Mapped[date] = mapped_column(Date, index=True) # 成交日期
direction: Mapped[str] = mapped_column(String(4)) # buy | sell
price: Mapped[float | None] = mapped_column(Float) # 成交价
qty: Mapped[float] = mapped_column(Float) # 股数
amount: Mapped[float | None] = mapped_column(Float) # 成交金额(元)
fee: Mapped[float] = mapped_column(Float, default=0.0) # 手续费合计(元)
raw_json: Mapped[str | None] = mapped_column(Text) # 原始行(审计/排错)
created_at: Mapped[datetime] = mapped_column(DateTime, default=_utcnow)
__table_args__ = (
# 重复上传同一份交割单幂等price 可空导致 PG 对 NULL 不去重,导入时另有 Python 侧兜底)
UniqueConstraint("user_id", "trade_date", "ts_code", "direction", "price", "qty", name="uq_user_trade_dedup"),
)
class ScreenerQuery(Base):
"""自然语言选股提问历史(文本 + 解析出的条件,便于一键重跑)。"""
__tablename__ = "screener_queries"

View File

@@ -4,7 +4,7 @@
"""
from __future__ import annotations
from datetime import datetime
from datetime import date, datetime
from typing import Literal
from pydantic import BaseModel, Field
@@ -298,7 +298,11 @@ class StockListItemOut(BaseModel):
prev_close: float | None = None
pct_chg: float | None = None # 最新两根日线计算
last_ts: datetime | None = None
bar_count: int | None = None # 本地缓存日线条数
turnover_rate: float | None = None # 换手率 %daily_snapshot
pe_ttm: float | None = None # 市盈率 TTM
pb: float | None = None # 市净率
total_mv: float | None = None # 总市值(亿元)
circ_mv: float | None = None # 流通市值(亿元)
watched: bool = False # 是否自选(当前用户)
@@ -343,3 +347,29 @@ class ScreenerQueryOut(BaseModel):
class ScreenerQueryListResponse(BaseModel):
items: list[ScreenerQueryOut]
# ---------- 交割单(个人实盘买卖点) ----------
class UserTradeOut(BaseModel):
id: int
ts_code: str
name: str | None = None
trade_date: date # 成交日期ISO YYYY-MM-DD
direction: str # buy | sell
price: float | None = None # 成交价(券商原始价,不复权)
qty: float # 股数
amount: float | None = None
fee: float | None = None
class TradesImportResponse(BaseModel):
inserted: int # 新入库成交笔数
skipped_dup: int # 与库内完全一致(重复上传同文件)跳过
skipped_other: int # 非买卖行(转账/配号/利息等)
stocks: int # 涉及股票数
bad: list[str] = Field(default_factory=list) # 解析失败样例(前 5 条)
sample: list[UserTradeOut] = Field(default_factory=list) # 本次入库的前几笔(核对用)
class TradesClearResponse(BaseModel):
deleted: int

329
backend/app/trades.py Normal file
View File

@@ -0,0 +1,329 @@
"""交割单解析(券商导出的成交流水 → 结构化买卖记录)。
支持三类导出物(按内容嗅探,不信任扩展名):
- CSV/制表符文本utf-8-sig / gbk / gb18030 自动探测)
- Excel .xlsxopenpyxl很多券商导出的 .xls 实为 xlsx 或 HTML先按魔数分流
- HTML 表格(.xls 常见真身:<table><tr><td>
列名模糊匹配兼容通达信/恒生/同花顺系的命名差异;业务名称含「买入/卖出」
才入库,银行转账、配号、利息、红利等非交易行跳过并计数。
"""
from __future__ import annotations
import csv
import io
import re
from dataclasses import dataclass, field
from datetime import date, datetime
from fastapi import HTTPException
@dataclass
class ParsedTrade:
trade_date: date
ts_code: str
code: str
name: str
direction: str # buy | sell
price: float | None
qty: float
amount: float | None
fee: float
raw: dict = field(default_factory=dict)
@dataclass
class ParseResult:
trades: list[ParsedTrade] = field(default_factory=list)
skipped_other: int = 0 # 非证券买卖行(转账/配号/利息等)
skipped_bad: list[str] = field(default_factory=list) # 解析失败样例(截断到前 5 条)
header_row_index: int = -1
columns: dict[str, str] = field(default_factory=dict) # 逻辑列 -> 实际列名
# ---------- 列名别名(归一化后做「包含」匹配,先命中的优先) ----------
COLUMN_ALIASES: dict[str, list[str]] = {
"date": ["成交日期", "交割日期", "交收日期", "交易日期", "过户日期", "发生日期", "清算日期", "日期"],
"op": ["业务名称", "业务摘要", "操作", "业务类型", "交易类型", "交易类别", "摘要", "方向", "买卖标志"],
"code": ["证券代码", "股票代码", "产品代码", "代码"],
"name": ["证券名称", "股票名称", "产品名称", "名称"],
"qty": ["成交数量", "发生数量", "委托数量", "成交股数", "数量"],
"price": ["成交价格", "成交均价", "成交价", "均价", "价格"],
"amount": ["成交金额", "成交清算金额", "清算金额", "发生金额", "资金发生数", "金额"],
"fee": ["手续费", "佣金", "印花税", "过户费", "其他费", "杂费", "规费"],
}
# 手续费类允许多列求和(手续费+印花税+过户费…),其余逻辑列取第一命中
_FEE_KEYS = ("手续费", "佣金", "印花税", "过户费", "其他费", "杂费", "规费")
def _norm_header(h: str) -> str:
"""列名归一化:去空白、去全角、去括号单位(如「成交数量(股)」)。"""
h = str(h).strip().replace(" ", "").replace(" ", "").replace(" ", "")
h = re.sub(r"[(【\[].*?[))】\]]", "", h)
return h
def _match_columns(header: list[str]) -> dict[str, str]:
"""表头 -> 逻辑列映射。返回 {逻辑列: 实际列名};费率类列全部收集到 fee(合并名)。"""
out: dict[str, str] = {}
fee_cols: list[str] = []
for h in header:
n = _norm_header(h)
if not n:
continue
for key, aliases in COLUMN_ALIASES.items():
if key == "fee":
if any(a in n for a in _FEE_KEYS):
fee_cols.append(h)
continue
if key in out:
continue
if any(a in n for a in aliases):
out[key] = h
break
# 「费用合计」列本身已含全部费用明细,取它即可,避免与手续费/印花税等列重复累加
total_col = next((h for h in header if "费用合计" in _norm_header(h)), None)
if total_col is not None:
out["fee"] = total_col
elif fee_cols:
out["fee"] = "\x00".join(fee_cols) # 多列合并存储,取值时拆开求和
return out
def _looks_like_header(row: list[str]) -> bool:
"""前 10 行里找表头≥3 个逻辑列可识别即认为是表头。"""
return len(_match_columns(row)) >= 3
def _to_float(v) -> float | None:
"""'1,234.50' / '(123.45)' / '--' / '' → float不可解析返回 None。"""
if v is None:
return None
if isinstance(v, (int, float)):
return float(v)
s = str(v).strip().replace(",", "").replace("", "")
if not s or s in {"--", "-", ""}:
return None
neg = s.startswith("(") and s.endswith(")")
if neg:
s = s[1:-1]
try:
f = float(s)
except ValueError:
return None
return -f if neg else f
def _to_date(v) -> date | None:
if isinstance(v, datetime):
return v.date()
if isinstance(v, date):
return v
if isinstance(v, (int, float)) and not isinstance(v, bool) and 30000 < v < 60000:
# Excel 日期序列值1982~2064openpyxl 读无日期格式的单元格时会给出
from datetime import timedelta
return date(1899, 12, 30) + timedelta(days=int(v))
s = str(v).strip()
m = re.search(r"(\d{4})[-/.年](\d{1,2})[-/.月](\d{1,2})", s)
if not m:
m2 = re.fullmatch(r"(\d{4})(\d{2})(\d{2})", s)
if not m2:
return None
m = m2
y, mo, d = int(m.group(1)), int(m.group(2)), int(m.group(3))
try:
return date(y, mo, d)
except ValueError:
return None
def _to_code_suffix(code: str) -> str:
"""6 位代码 → 交易所后缀60/68 沪00/30 深4/8/92 北交所)。"""
if code.startswith(("60", "68", "90")):
return ".SH"
if code.startswith(("00", "30", "20")):
return ".SZ"
return ".BJ"
def _direction(op: str) -> str | None:
s = str(op)
if "买入" in s or "buy" in s.lower() or "证券买" in s:
return "buy"
if "卖出" in s or "sell" in s.lower() or "证券卖" in s:
return "sell"
return None
def _parse_rows(rows: list[list[object]]) -> ParseResult:
"""已抽成二维表的行集 → ParseResult。rows[0] 应是表头(调用方已定位)。"""
res = ParseResult()
if not rows:
return res
header = [str(h) for h in rows[0]]
cols = _match_columns(header)
res.columns = {k: v for k, v in cols.items()}
res.header_row_index = 0
need = ("date", "qty")
if not all(k in cols for k in need) or not ("code" in cols or "name" in cols):
raise HTTPException(
status_code=422,
detail="识别不到交割单表头(需要 成交日期/证券代码或证券名称/成交数量 等列),"
"请确认导出的是「交割单/历史成交」文件",
)
idx = {h: i for i, h in enumerate(header)}
# 无「业务名称」列的导出(如部分招商证券格式):靠发生金额正负判方向(买入为负)。
# 仅当数据里确实存在负数金额才启用,避免「全正数」格式被误判。
def _amount_of(row: list[object]) -> float | None:
i = idx.get(cols["amount"])
return _to_float(row[i]) if i is not None and i < len(row) else None
sign_mode = "op" not in cols and "amount" in cols and any(
(_amount_of(row) or 0) < 0 for row in rows[1:] if any(str(c).strip() for c in row)
)
def cell(row: list[object], col: str):
i = idx.get(col)
return row[i] if i is not None and i < len(row) else None
for row in rows[1:]:
d = _to_date(cell(row, cols["date"]))
code = re.sub(r"\D", "", str(cell(row, cols["code"]) or "")) if "code" in cols else ""
raw_amount = _amount_of(row) if sign_mode else None
direction = (
_direction(str(cell(row, cols["op"]) or "")) if "op" in cols
else ("buy" if (raw_amount or 0) < 0 else "sell") if sign_mode
else None
)
name = str(cell(row, cols["name"]) or "").strip() if "name" in cols else ""
if d is None or (not code and not name) or direction is None:
# 无日期/无代码且无名称/非买卖业务(银行转账、配号、利息、红利等)
if any(str(c).strip() for c in row):
res.skipped_other += 1
continue
if len(code) > 6:
code = code[-6:] # 个别导出带市场前缀(如 1:600000 / sh600000
qty = abs(_to_float(cell(row, cols["qty"])) or 0)
if qty <= 0:
res.skipped_bad.append(f"{d} {code or name} 数量无效:{cell(row, cols['qty'])!r}")
continue
price = _to_float(cell(row, cols["price"])) if "price" in cols else None
amount = raw_amount if sign_mode else (_to_float(cell(row, cols["amount"])) if "amount" in cols else None)
if amount is not None:
amount = abs(amount)
fee = 0.0
if "fee" in cols:
for fc in cols["fee"].split("\x00"):
f = _to_float(cell(row, fc))
if f:
fee += abs(f)
# 无代码列招商式导出ts_code 留空,由 API 层按 name 反查 stock_basic
ts_code = code + _to_code_suffix(code) if code else ""
res.trades.append(ParsedTrade(
trade_date=d,
code=code,
ts_code=ts_code,
name=name,
direction=direction,
price=price,
qty=qty,
amount=amount,
fee=round(fee, 2),
raw={h: row[i] if i < len(row) else None for i, h in enumerate(header)},
))
res.skipped_bad = res.skipped_bad[:5]
return res
def _find_header(rows: list[list[object]]) -> int:
for i, row in enumerate(rows[:10]):
if _looks_like_header([str(c) for c in row]):
return i
return -1
# ---------- 输入格式分流 ----------
def _rows_from_csv(data: bytes) -> list[list[object]]:
"""逗号/制表符分隔文本。sniff 分隔符;跳过全空行。"""
text = None
for enc in ("utf-8-sig", "gbk", "gb18030"):
try:
text = data.decode(enc)
break
except UnicodeDecodeError:
continue
if text is None:
raise HTTPException(status_code=422, detail="文件编码无法识别(支持 UTF-8 / GBK")
sample = text[:4096]
delim = "\t" if sample.count("\t") > sample.count(",") else ","
lines = [ln for ln in text.splitlines() if ln.strip()]
if not lines:
raise HTTPException(status_code=422, detail="文件是空的")
return [next(csv.reader([ln], delimiter=delim)) for ln in lines]
def _rows_from_xlsx(data: bytes) -> list[list[object]]:
from openpyxl import load_workbook
try:
wb = load_workbook(io.BytesIO(data), read_only=True, data_only=True)
except Exception as e: # noqa: BLE001 - openpyxl 对损坏文件抛各种类型
raise HTTPException(status_code=422, detail=f"Excel 文件无法读取:{e}") from e
ws = wb.active
rows = [[c for c in row] for row in ws.iter_rows(values_only=True)]
wb.close()
return rows
_TD_RE = re.compile(r"<t[dh][^>]*>(.*?)</t[dh]>", re.IGNORECASE | re.DOTALL)
_TR_RE = re.compile(r"<tr[^>]*>(.*?)</tr>", re.IGNORECASE | re.DOTALL)
def _rows_from_html(data: bytes) -> list[list[object]]:
"""券商导出的 .xls 常是 HTML 表格。去掉标签实体后按 <tr>/<td> 切。"""
text = None
for enc in ("utf-8", "gbk", "gb18030"):
try:
text = data.decode(enc)
break
except UnicodeDecodeError:
continue
if text is None:
raise HTTPException(status_code=422, detail="文件编码无法识别(支持 UTF-8 / GBK")
import html as html_mod
rows: list[list[object]] = []
for tr in _TR_RE.findall(text):
cells = [html_mod.unescape(re.sub(r"<[^>]+>", "", td)).strip() for td in _TD_RE.findall(tr)]
rows.append(cells)
if not rows:
raise HTTPException(status_code=422, detail="HTML 里没有表格数据")
return rows
def parse_statement(data: bytes, filename: str) -> ParseResult:
"""入口:按内容魔数/特征分流 → 定位表头 → 解析。"""
if not data:
raise HTTPException(status_code=422, detail="文件是空的")
head = data[:512].lstrip()
if head.startswith(b"PK"):
rows = _rows_from_xlsx(data)
elif head[:1] in (b"<",) or head.lower().startswith(b"\xef\xbb\xbf<"):
rows = _rows_from_html(data)
elif filename.lower().endswith((".xlsx", ".xls")) and not head.startswith((b"PK", b"<")):
# 扩展名是 Excel 但内容既非 xlsx 也非 HTML → 试试当文本
rows = _rows_from_csv(data)
else:
rows = _rows_from_csv(data)
# 去尾部全空行,定位表头(导出物常有标题行/账户信息行在前)
while rows and not any(str(c).strip() for c in rows[-1]):
rows.pop()
hi = _find_header(rows)
if hi < 0:
raise HTTPException(
status_code=422,
detail="找不到表头行(前 10 行内没有 成交日期/证券代码 等列名),请确认导出的是交割单",
)
return _parse_rows(rows[hi:])