看股功能更新
This commit is contained in:
@@ -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)
|
||||
|
||||
265
backend/app/backtest/events.py
Normal file
265
backend/app/backtest/events.py
Normal file
@@ -0,0 +1,265 @@
|
||||
"""事件回测引擎:入场条件命中 -> 次日买入 -> 持有 N 日 -> 全市场汇总统计。
|
||||
|
||||
数据口径:
|
||||
- 行情底座是 candles(TDX 全量导入,不复权),全历史可用;
|
||||
- 指标计算用不复权价(与选股/看盘口径一致: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 lookback)AND。"""
|
||||
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
104
backend/app/cache.py
Normal 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)。版本号参与缓存 key:INCR 后旧 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
|
||||
@@ -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(智能选股的自然语言解析;DeepSeek,OpenAI 兼容协议,可换任意兼容网关)----
|
||||
llm_base_url: str = "https://api.deepseek.com"
|
||||
llm_api_key: str = "" # 留空则智能选股不可用(其余功能不受影响)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
329
backend/app/trades.py
Normal file
@@ -0,0 +1,329 @@
|
||||
"""交割单解析(券商导出的成交流水 → 结构化买卖记录)。
|
||||
|
||||
支持三类导出物(按内容嗅探,不信任扩展名):
|
||||
- CSV/制表符文本(utf-8-sig / gbk / gb18030 自动探测)
|
||||
- Excel .xlsx(openpyxl;很多券商导出的 .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~2064),openpyxl 读无日期格式的单元格时会给出
|
||||
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:])
|
||||
Reference in New Issue
Block a user