看股功能更新

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

@@ -3,6 +3,9 @@ TUSHARE_TOKEN=22edda0afe44c0609a187ff1ac0bb2a8fc61430f490ec19f7fec8390
DATA_ADJUST=qfq
DATA_DEFAULT_START=20200101
# ---- Redis 读缓存(股票列表/筛选项;留空则不缓存直查数据库)----
REDIS_URL=redis://default:26d5c71d57344f37b8b4ddb567f2652f0c7ef41c774284ad@cirry.cn:6379
# ---- LLM智能选股智谱 GLMOpenAI 兼容协议)----
# key 在 https://bigmodel.cn 控制台获取,格式形如 xxxxxxxx.yyyyyyyyid.secret
LLM_BASE_URL=https://open.bigmodel.cn/api/paas/v4

View File

@@ -38,3 +38,9 @@ LLM_MODEL=glm-5.2
# TRANSFER_FEE_RATE=0.00001 # 过户费 0.001%,沪深双边
# COMMISSION_RATE=0.0001 # 佣金 万1
# COMMISSION_MIN=5.0 # 最低 5 元
# ---- Redis 读缓存(可选;股票列表/筛选项提速)----
# 留空 = 不缓存,直查数据库;连接失败自动降级,不影响接口可用性
# REDIS_URL=redis://default:password@127.0.0.1:6379
# STOCKS_CACHE_TTL=300 # 股票列表缓存秒数
# FACETS_CACHE_TTL=3600 # 行业/地域筛选项缓存秒数

View File

@@ -0,0 +1,36 @@
"""add adj_factor table (复权因子底座)
Revision ID: 20260815_01
Revises: 208b0c5d302a
Create Date: 2026-08-15
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "20260815_01"
down_revision: Union[str, Sequence[str], None] = "208b0c5d302a"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.create_table(
"adj_factor",
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
sa.Column("trade_date", sa.DateTime(), nullable=False),
sa.Column("ts_code", sa.String(length=12), nullable=False),
sa.Column("adj_factor", sa.Float(), nullable=False),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("ts_code", "trade_date", name="uq_adj_code_date"),
)
op.create_index("ix_adj_factor_trade_date", "adj_factor", ["trade_date"])
op.create_index("ix_adj_factor_ts_code", "adj_factor", ["ts_code"])
def downgrade() -> None:
op.drop_index("ix_adj_factor_ts_code", table_name="adj_factor")
op.drop_index("ix_adj_factor_trade_date", table_name="adj_factor")
op.drop_table("adj_factor")

View File

@@ -0,0 +1,69 @@
"""user data tables: preferences / watchlist / screener queries
Revision ID: 20260815_02
Revises: 20260815_01
Create Date: 2026-08-15
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "20260815_02"
down_revision: Union[str, Sequence[str], None] = "20260815_01"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.create_table(
"user_preferences",
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
sa.Column("user_id", sa.BigInteger(), nullable=False),
sa.Column("key", sa.String(length=64), nullable=False),
sa.Column("value_json", sa.Text(), nullable=False, server_default="null"),
sa.Column("updated_at", sa.DateTime(), nullable=False),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("user_id", "key", name="uq_user_pref_key"),
)
op.create_index("ix_user_preferences_user_id", "user_preferences", ["user_id"])
op.create_table(
"watchlist_items",
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
sa.Column("user_id", sa.BigInteger(), nullable=False),
sa.Column("ts_code", sa.String(length=12), nullable=False),
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("user_id", "ts_code", name="uq_watch_user_code"),
)
op.create_index("ix_watchlist_items_user_id", "watchlist_items", ["user_id"])
op.create_index("ix_watchlist_items_ts_code", "watchlist_items", ["ts_code"])
op.create_table(
"screener_queries",
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
sa.Column("user_id", sa.BigInteger(), nullable=False),
sa.Column("text", sa.String(length=500), nullable=False),
sa.Column("conditions_json", sa.Text(), nullable=True),
sa.Column("hit_count", sa.Integer(), nullable=True),
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_screener_queries_user_id", "screener_queries", ["user_id"])
op.create_index("ix_screener_queries_created_at", "screener_queries", ["created_at"])
def downgrade() -> None:
op.drop_index("ix_screener_queries_created_at", table_name="screener_queries")
op.drop_index("ix_screener_queries_user_id", table_name="screener_queries")
op.drop_table("screener_queries")
op.drop_index("ix_watchlist_items_ts_code", table_name="watchlist_items")
op.drop_index("ix_watchlist_items_user_id", table_name="watchlist_items")
op.drop_table("watchlist_items")
op.drop_index("ix_user_preferences_user_id", table_name="user_preferences")
op.drop_table("user_preferences")

View File

@@ -0,0 +1,30 @@
"""candles add amount/turnover columns
Revision ID: 20260815_03
Revises: 20260815_02
Create Date: 2026-08-15
- amount 成交额TDX .day 原生 float32/ Tushare daily amount 千元×1000
- turnover 换手率(%Tushare daily_basic.turnover_rate2000 年起)
均为可空列——历史回补前为 NULL前端 tooltip 显示 ""
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "20260815_03"
down_revision: Union[str, Sequence[str], None] = "20260815_02"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.add_column("candles", sa.Column("amount", sa.Float(), nullable=True))
op.add_column("candles", sa.Column("turnover", sa.Float(), nullable=True))
def downgrade() -> None:
op.drop_column("candles", "turnover")
op.drop_column("candles", "amount")

View File

@@ -0,0 +1,51 @@
"""user_trades: 交割单导入的实盘成交流水K线买卖点数据源
Revision ID: 20260815_04
Revises: 20260815_03
Create Date: 2026-08-15
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "20260815_04"
down_revision: Union[str, Sequence[str], None] = "20260815_03"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.create_table(
"user_trades",
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
sa.Column("user_id", sa.BigInteger(), nullable=False),
sa.Column("ts_code", sa.String(length=12), nullable=False),
sa.Column("code", sa.String(length=10), nullable=False),
sa.Column("name", sa.String(length=32), nullable=True),
sa.Column("trade_date", sa.Date(), nullable=False),
sa.Column("direction", sa.String(length=4), nullable=False),
sa.Column("price", sa.Float(), nullable=True),
sa.Column("qty", sa.Float(), nullable=False),
sa.Column("amount", sa.Float(), nullable=True),
sa.Column("fee", sa.Float(), nullable=False, server_default="0"),
sa.Column("raw_json", sa.Text(), nullable=True),
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint(
"user_id", "trade_date", "ts_code", "direction", "price", "qty",
name="uq_user_trade_dedup",
),
)
op.create_index("ix_user_trades_user_id", "user_trades", ["user_id"])
op.create_index("ix_user_trades_ts_code", "user_trades", ["ts_code"])
op.create_index("ix_user_trades_trade_date", "user_trades", ["trade_date"])
def downgrade() -> None:
op.drop_index("ix_user_trades_trade_date", table_name="user_trades")
op.drop_index("ix_user_trades_ts_code", table_name="user_trades")
op.drop_index("ix_user_trades_user_id", table_name="user_trades")
op.drop_table("user_trades")

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:])

View File

@@ -16,6 +16,9 @@ dependencies = [
"httpx>=0.28.1",
"argon2-cffi>=25.1.0",
"alembic>=1.19.1",
"redis>=8.1.0",
"python-multipart>=0.0.32",
"openpyxl>=3.1.5",
]
[tool.uv]

View File

@@ -0,0 +1,122 @@
"""全量回补历史复权因子adj_factor 表)。
用法(在 backend 目录下):
uv run python scripts/backfill_adj_factor.py # 从 candles 最早日期回补到今天
uv run python scripts/backfill_adj_factor.py --start 20180101
uv run python scripts/backfill_adj_factor.py --force # 已有日期也重拉
- 按交易日逐日拉取全市场因子pro.adj_factor(trade_date=...)),幂等可断点续跑;
- 交易日取自本地 trade_calendar缓存不到的区间自动刷新一次日历
- Tushare 每分钟限频由 _call_retry 自动等待 62s 重试。
"""
from __future__ import annotations
import argparse
import asyncio
import sys
import time
from datetime import datetime
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sqlalchemy import delete, func, insert, select
from app.db import async_session
from app.models import AdjFactor, Candle, TradeCalendar
from app.screener.market_sync import _call_retry, _get_pro, _norm_date, _parse_d
_INTERVAL_MSG = 20 # 每完成 N 个交易日打印一次进度
async def _calendar_dates(start: str, end: str) -> list[str]:
"""[start, end] 交易日(升序)。本地日历覆盖不足时直接拉宽范围日历并回写缓存。"""
async with async_session() as session:
all_cached = set((await session.execute(select(TradeCalendar.trade_date))).scalars().all())
cached = sorted(d for d in all_cached if start <= d <= end)
if cached and min(cached) <= start:
return cached
# 覆盖不到起点按需拉宽范围日历trade_cal 低积分限频 1 次/小时,失败沿用缓存)
pro = _get_pro()
try:
cal = await asyncio.to_thread(
_call_retry, pro.trade_cal, exchange="SSE", start_date=start, end_date=end, is_open="1"
)
dates = sorted(cal["cal_date"].tolist())
except Exception as e: # noqa: BLE001
if not cached:
raise
print(f"交易日历拉取受限({str(e)[:100]}),沿用本地缓存")
return cached
fresh = [d for d in dates if d not in all_cached]
if fresh:
async with async_session() as session:
await session.execute(insert(TradeCalendar), [{"trade_date": d} for d in fresh])
await session.commit()
return dates
async def _existing_dates() -> set[str]:
async with async_session() as session:
res = await session.execute(select(func.distinct(AdjFactor.trade_date)))
return {_norm_date(r[0]) for r in res}
async def main(start: str, end: str, force: bool) -> None:
# 默认起点candles 最早日线(因子只需覆盖有 K 线的区间)
if start is None:
async with async_session() as session:
first = await session.scalar(select(func.min(Candle.ts)).where(Candle.timeframe == "1d"))
start = first.strftime("%Y%m%d") if first else "20050101"
if end is None:
end = datetime.now().strftime("%Y%m%d")
dates = await _calendar_dates(start, end)
have = set() if force else await _existing_dates()
todo = [d for d in dates if d not in have]
print(f"区间 {start}~{end}{len(dates)} 个交易日,待回补 {len(todo)} 个(已有 {len(dates) - len(todo)}")
if not todo:
return
pro = _get_pro()
done = 0
for d in todo:
time.sleep(0.15) # 轻微控频;分钟级限频由 _call_retry 自动等待重试
df = None
for attempt in range(5): # 网络抖动(超时/断连也重试_call_retry 只兜限频
try:
df = _call_retry(pro.adj_factor, trade_date=d) # noqa: 线性脚本直接同步调用
break
except Exception as e: # noqa: BLE001
wait = min(30 * (attempt + 1), 120)
print(f" {d} 拉取异常({str(e)[:80]}{wait}s 后重试 {attempt + 1}/5")
time.sleep(wait)
if df is None:
print(f" {d} 连续 5 次失败,跳过(断点续跑可补)")
continue
if df is None or df.empty:
print(f" {d} 无数据(非交易日或未生成),跳过")
continue
rows = [
{"trade_date": _parse_d(d), "ts_code": r["ts_code"], "adj_factor": float(r["adj_factor"])}
for _, r in df.iterrows()
]
async with async_session() as session:
dt = _parse_d(d)
await session.execute(delete(AdjFactor).where(AdjFactor.trade_date == dt))
await session.execute(insert(AdjFactor), rows)
await session.commit()
done += 1
if done % _INTERVAL_MSG == 0 or done == len(todo):
print(f" 进度 {done}/{len(todo)}{d}+{len(rows)} 行)")
print(f"回补完成:{done} 个交易日")
if __name__ == "__main__":
ap = argparse.ArgumentParser(description="全量回补历史复权因子")
ap.add_argument("--start", default=None, help="YYYYMMDD默认 candles 最早日期")
ap.add_argument("--end", default=None, help="YYYYMMDD默认今天")
ap.add_argument("--force", action="store_true", help="已有日期也重拉")
a = ap.parse_args()
asyncio.run(main(a.start, a.end, a.force))

View File

@@ -0,0 +1,162 @@
"""全量回补换手率candles.turnover单位 %)。
用法(在 backend 目录下):
uv run python scripts/backfill_turnover.py # 从 2000-01-01daily_basic 起点)回补到今天
uv run python scripts/backfill_turnover.py --start 20200101
uv run python scripts/backfill_turnover.py --force # 已回补的交易日也重拉
- 数据源Tushare daily_basic(trade_date=..., fields='ts_code,turnover_rate'),按日全市场;
- 幂等可断点续跑:某交易日 candles 已有非空 turnover 即跳过(--force 强制重做);
- 交易日取自本地 trade_calendar缓存覆盖不到起点时自动拉一次宽范围日历
- 每日一条 UPDATE ... FROM unnest(...) 批量写回,仅更新 turnover 列;
- Tushare 每分钟限频由 _call_retry 自动等待 62s 重试。
注意:与 import_tdx_day.py回填 amount 会整行 upsert串行运行避免同表行锁竞争。
"""
from __future__ import annotations
import argparse
import asyncio
import sys
import time
from datetime import datetime
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from app.screener.market_sync import _call_retry, _get_pro
import asyncpg
def load_db_url() -> str:
"""与 import_tdx_day.py 相同的 .env -> libpq URL 解析(本地复制避免跨脚本导入)。"""
env = Path(__file__).resolve().parent.parent / ".env"
if env.exists():
for line in env.read_text(encoding="utf-8").splitlines():
line = line.strip()
if line.startswith("DATABASE_URL=postgresql+asyncpg://"):
return "postgresql://" + line.split("://", 1)[1]
return "postgresql://postgres:postgres@localhost:5432/stock"
_DAILY_BASIC_FLOOR = "20000101" # daily_basic 最早覆盖 2000-01-04更早的交易日无换手数据
_INTERVAL_MSG = 20
async def _calendar_dates(conn: asyncpg.Connection, start: str, end: str) -> list[str]:
"""[start, end] 交易日(升序)。本地缓存覆盖不到起点时拉一次宽范围日历并回写。"""
cached = [r[0] for r in await conn.fetch(
"SELECT trade_date FROM trade_calendar WHERE trade_date >= $1 AND trade_date <= $2 "
"ORDER BY trade_date", start, end)]
if cached and cached[0] <= start:
return cached
pro = _get_pro()
try:
cal = await asyncio.to_thread(
_call_retry, pro.trade_cal, exchange="SSE", start_date=start, end_date=end, is_open="1"
)
dates = sorted(cal["cal_date"].tolist())
except Exception as e: # noqa: BLE001
if not cached:
raise
print(f"交易日历拉取受限({str(e)[:100]}),沿用本地缓存")
return cached
have = set(cached)
fresh = [d for d in dates if d not in have]
if fresh:
await conn.executemany(
"INSERT INTO trade_calendar (trade_date) VALUES ($1) ON CONFLICT DO NOTHING", [(d,) for d in fresh]
)
return dates
async def _day_status(conn: asyncpg.Connection, d: str) -> tuple[int, int]:
"""(已有换手的行数, 当日总行数)。无行情的日子 total=0 直接跳过。"""
row = await conn.fetchrow(
"SELECT count(*) FILTER (WHERE turnover IS NOT NULL) AS done, count(*) AS total "
"FROM candles WHERE timeframe = '1d' AND ts = $1::timestamp", datetime.strptime(d, "%Y%m%d")
)
return row["done"], row["total"]
async def main(start: str, end: str, force: bool) -> None:
conn = await asyncpg.connect(load_db_url())
try:
# 默认起点daily_basic 覆盖范围与 candles 最早日线的较大者(更早的日期拉了也是空)
if start is None:
first = await conn.fetchval(
"SELECT min(ts) FROM candles WHERE timeframe = '1d' AND symbol <> 'DEMO'")
start = max(first.strftime("%Y%m%d"), _DAILY_BASIC_FLOOR) if first else _DAILY_BASIC_FLOOR
if end is None:
end = datetime.now().strftime("%Y%m%d")
dates = await _calendar_dates(conn, start, end)
todo: list[str] = []
for d in dates:
if force:
done, total = await _day_status(conn, d)
if total:
todo.append(d)
continue
done, total = await _day_status(conn, d)
if total and done < total // 2: # 过半缺换手才重做(容忍个别股票无快照)
todo.append(d)
print(f"区间 {start}~{end}{len(dates)} 个交易日,待回补 {len(todo)}")
pro = _get_pro()
done = 0
t0 = time.time()
for d in todo:
time.sleep(0.15) # 轻微控频;分钟级限频由 _call_retry 自动等待重试
df = None
for attempt in range(5): # 网络抖动(超时/断连也重试_call_retry 只兜限频
try:
df = _call_retry(
pro.daily_basic, trade_date=d, fields="ts_code,trade_date,turnover_rate"
)
break
except Exception as e: # noqa: BLE001
wait = min(30 * (attempt + 1), 120)
print(f" {d} 拉取异常({str(e)[:80]}{wait}s 后重试 {attempt + 1}/5")
time.sleep(wait)
if df is None:
print(f" {d} 连续 5 次失败,跳过(断点续跑可补)")
continue
if df.empty:
continue
syms: list[str] = []
vals: list[float] = []
for _, r in df.iterrows():
tr = r["turnover_rate"]
if tr is None or tr != tr: # None / NaN
continue
syms.append(str(r["ts_code"]).split(".")[0])
vals.append(float(tr))
if not syms:
continue
n = await conn.execute(
"UPDATE candles AS c SET turnover = v.t "
"FROM unnest($1::text[], $2::float8[]) AS v(sym, t) "
"WHERE c.symbol = v.sym AND c.timeframe = '1d' AND c.ts = $3::timestamp",
syms, vals, datetime.strptime(d, "%Y%m%d"),
)
done += 1
if done % _INTERVAL_MSG == 0 or done == len(todo):
elapsed = time.time() - t0
eta = elapsed / done * (len(todo) - done) if done else 0
print(f" 进度 {done}/{len(todo)}{d}{len(syms)} 只,{n}"
f"{elapsed:.0f}s 已用,预计还需 {eta/60:.0f}m")
print(f"回补完成:{done} 个交易日")
finally:
await conn.close()
if __name__ == "__main__":
ap = argparse.ArgumentParser(description="全量回补换手率 candles.turnover")
ap.add_argument("--start", default=None, help="YYYYMMDD默认 max(candles 最早, 20000101)")
ap.add_argument("--end", default=None, help="YYYYMMDD默认今天")
ap.add_argument("--force", action="store_true", help="已有换手的交易日也重拉")
a = ap.parse_args()
asyncio.run(main(a.start, a.end, a.force))

View File

@@ -0,0 +1,160 @@
"""通达信「沪深京日线数据完整包」全量导入 candles 表。
用法(在 backend 目录下):
uv run python scripts/import_tdx_day.py C:/Users/cirry/Downloads/hsjday [symbol ...]
# symbol 为可选的 6 位代码过滤(如 000001 002671只重导这些标的
uv run python scripts/import_tdx_day.py <目录> --no-clear
# --no-clear不清空任何行纯 upsert用于给已导入的底座回补 amount 成交额)
- 解析 vipdoc 的 .day 二进制文件(每条 32 字节):
日期(YYYYMMDD) 开 高 低 收×100 成交额(元, float32) 成交量(股) 保留
- 只导入 stock_basic 里登记的股票(自动排除指数/基金/可转债/回购);
sh000001(上证指数) 与 sz000001(平安银行) 这类代码冲突也由此化解。
- 价格为**不复权**:全量模式导入前清空已有的非 DEMO 行情;指定 symbol 过滤时
只清空这些标的(用于修复被复权口径污染的个别股票),其余不动。
- amount 为 TDX 原生 float32精度 ~6 位有效数字,展示用途足够;
ON CONFLICT 时仅更新 amount 列,不动 OHLCV/turnover避免与换手率回补互相干扰
- 写入用 asyncpg execute_many + ON CONFLICT DO UPDATE可重复执行幂等
"""
from __future__ import annotations
import argparse
import asyncio
import struct
import sys
import time
from pathlib import Path
import asyncpg
# .env 里的 DATABASE_URL 是 SQLAlchemy 格式asyncpg 需要 libpq 格式
DEFAULT_URL = "postgresql://postgres:postgres@localhost:5432/stock"
BATCH = 20_000 # 每批 upsert 行数
def load_db_url() -> str:
env = Path(__file__).resolve().parent.parent / ".env"
if env.exists():
for line in env.read_text(encoding="utf-8").splitlines():
line = line.strip()
if line.startswith("DATABASE_URL=postgresql+asyncpg://"):
return "postgresql://" + line.split("://", 1)[1]
return DEFAULT_URL
def parse_day_file(path: Path) -> list[tuple[int, float, float, float, float, float, float]]:
"""解析单个 .day 文件 -> [(date, open, high, low, close, volume(股), amount(元)), ...]"""
raw = path.read_bytes()
unpack = struct.Struct("<IIIIIfII").unpack_from
out = []
for i in range(len(raw) // 32):
date, o, h, l, c, amount, vol, _reserved = unpack(raw, i * 32)
out.append((date, o / 100.0, h / 100.0, l / 100.0, c / 100.0, float(vol), float(amount)))
return out
async def main(root: Path, symbols: list[str] | None = None, no_clear: bool = False) -> None:
if not root.exists():
sys.exit(f"目录不存在: {root}")
conn = await asyncpg.connect(load_db_url())
try:
# 股票清单ts_code 形如 000001.SZ用于过滤指数/基金/转债
rows = await conn.fetch("SELECT ts_code, symbol FROM stock_basic WHERE list_status = 'L'")
by_exchange: dict[str, set[str]] = {"sh": set(), "sz": set(), "bj": set()}
for r in rows:
suffix = r["ts_code"].split(".")[-1].lower() # SH/SZ/BJ -> sh/sz/bj
if suffix in by_exchange:
by_exchange[suffix].add(r["symbol"])
print(f"stock_basic 在市股票: " + ", ".join(f"{k}={len(v)}" for k, v in by_exchange.items()))
files = sorted(root.glob("*/lday/*.day"))
print(f"发现 .day 文件: {len(files)}")
if no_clear:
print("--no-clear不清空任何行纯 upsert 回补 amount")
elif symbols:
# 清空旧行情(保留 DEMO 合成数据),避免 qfq/不复权混用;
# 带 symbol 过滤时只清空目标标的(修复个别被污染的股票,不动其余底座)
deleted = await conn.execute(
"DELETE FROM candles WHERE symbol = ANY($1)", symbols
)
print(f"清空目标标的 {symbols}: {deleted}")
keep = set(symbols)
files = [p for p in files if p.name[2:8] in keep]
print(f"过滤后待导入 .day 文件: {len(files)}")
else:
deleted = await conn.execute("DELETE FROM candles WHERE symbol <> 'DEMO'")
print(f"清空旧行情: {deleted}")
if no_clear:
# 回填模式:只写 amount不动 OHLCV/turnover底座已就位避免全表重写
upsert_sql = """
INSERT INTO candles (symbol, timeframe, ts, open, high, low, close, volume, amount)
VALUES ($1, '1d', to_timestamp($2::text, 'YYYYMMDD')::timestamp, $3, $4, $5, $6, $7, $8)
ON CONFLICT (symbol, timeframe, ts) DO UPDATE
SET amount = EXCLUDED.amount
"""
else:
upsert_sql = """
INSERT INTO candles (symbol, timeframe, ts, open, high, low, close, volume, amount)
VALUES ($1, '1d', to_timestamp($2::text, 'YYYYMMDD')::timestamp, $3, $4, $5, $6, $7, $8)
ON CONFLICT (symbol, timeframe, ts) DO UPDATE
SET open = EXCLUDED.open, high = EXCLUDED.high, low = EXCLUDED.low,
close = EXCLUDED.close, volume = EXCLUDED.volume, amount = EXCLUDED.amount
"""
t0 = time.time()
total_stocks = 0
skipped = 0
batch: list[tuple] = []
rows_done = 0
async def flush() -> None:
nonlocal batch, rows_done
if batch:
await conn.executemany(upsert_sql, batch)
rows_done += len(batch)
batch = []
for n, path in enumerate(files, 1):
market = path.name[:2].lower() # sh / sz / bj
code = path.name[2:8]
if code not in by_exchange.get(market, set()):
skipped += 1
continue
for date, o, h, l, c, v, amount in parse_day_file(path):
batch.append((code, str(date), o, h, l, c, v, amount))
total_stocks += 1
if len(batch) >= BATCH:
await flush()
if n % 500 == 0:
elapsed = time.time() - t0
print(f" 进度 {n}/{len(files)} 文件, 已入库 {total_stocks} 只股票, "
f"{rows_done + len(batch):,} 行, {elapsed:.0f}s")
await flush()
cnt = await conn.fetchval("SELECT count(*) FROM candles WHERE symbol <> 'DEMO'")
span = await conn.fetchrow(
"SELECT min(ts) AS lo, max(ts) AS hi FROM candles WHERE symbol <> 'DEMO'"
)
with_amt = await conn.fetchval(
"SELECT count(*) FROM candles WHERE symbol <> 'DEMO' AND amount IS NOT NULL"
)
print(f"\n完成: {total_stocks} 只股票, {cnt:,} 行日线, "
f"范围 {span['lo']:%Y-%m-%d} ~ {span['hi']:%Y-%m-%d}, "
f"含成交额 {with_amt:,} 行, "
f"跳过非股票文件 {skipped} 个, 耗时 {time.time() - t0:.0f}s")
finally:
await conn.close()
if __name__ == "__main__":
ap = argparse.ArgumentParser(description="TDX 沪深京日线全量导入 candles")
ap.add_argument("root", help="hsjday 目录(其下 */lday/*.day")
ap.add_argument("symbols", nargs="*", help="可选的 6 位代码过滤")
ap.add_argument("--no-clear", action="store_true",
help="不清空任何行,纯 upsertamount 回补模式)")
a = ap.parse_args()
asyncio.run(main(Path(a.root), a.symbols or None, a.no_clear))

View File

@@ -0,0 +1,113 @@
"""交割单解析器离线自测:不碰数据库,直接调 app.trades.parse_statement。
覆盖四类真实导出格式 + 边界行(转账/配号/利息跳过、费用合计列去重、日期多格式)。
运行uv run python scripts/test_trades_parser.py
"""
from __future__ import annotations
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from app.trades import parse_statement # noqa: E402
FAIL: list[str] = []
def check(name: str, cond: bool, detail: str = "") -> None:
mark = "ok " if cond else "FAIL"
print(f"[{mark}] {name}{('' + detail) if detail and not cond else ''}")
if not cond:
FAIL.append(name)
# ---------- 1) 通达信式GBK + 制表符 + 标题行在前 ----------
tdx = (
"交割单\n"
"股东账号: A123456789 起始日期: 20240102 终止日期: 20240105 币种: 人民币\n"
"\t交割日期\t业务名称\t证券代码\t证券名称\t成交价格\t成交数量\t成交金额\t手续费\t印花税\t过户费\t发生金额\t资金余额\t合同号\n"
"\t20240102\t证券买入\t600519\t贵州茅台\t1680.00\t100\t168000.00\t5.00\t0.00\t1.68\t-168006.68\t200000.00\t1000001\n"
"\t20240102\t银行转存\t\t\t\t\t\t\t\t\t50000.00\t250000.00\t\n"
"\t20240103\t证券卖出\t600519\t贵州茅台\t1700.50\t100\t170050.00\t5.00\t170.05\t1.70\t169873.25\t419873.25\t1000002\n"
"\t20240105\t利息归本\t\t\t\t\t\t\t\t\t1.25\t419874.50\t\n"
)
r = parse_statement(tdx.encode("gbk"), "交割单.txt")
check("tdx: 2 笔成交", len(r.trades) == 2, f"got {len(r.trades)}")
check("tdx: 跳过 2 行非交易", r.skipped_other == 2, f"got {r.skipped_other}")
t0, t1 = r.trades[0], r.trades[1]
check("tdx: 日期/代码/后缀", (t0.trade_date.isoformat(), t0.ts_code) == ("2024-01-02", "600519.SH"), f"{t0.trade_date} {t0.ts_code}")
check("tdx: 买入方向+费用合计", t0.direction == "buy" and abs(t0.fee - 6.68) < 1e-9, f"{t0.direction} fee={t0.fee}")
check("tdx: 卖出费用含印花税", t1.direction == "sell" and abs(t1.fee - 176.75) < 1e-9, f"fee={t1.fee}")
check("tdx: 金额取绝对值", t0.amount == 168000.0, f"amount={t0.amount}")
# ---------- 2) 恒生柜台式UTF-8 CSV交收日期/交易类别/费用合计 ----------
hs = (
"序号,交收日期,证券代码,证券名称,交易类别,成交价格,成交数量,证券余额,成交金额,资金发生数,资金余额,流水序号,业务标志,业务名称,发生金额,后资金额,货币类别,费用合计,净佣金,规费,印花税,过户费,合同号\n"
"1,2024-06-07,000858,五粮液,证券买入,132.50,200,200,26500.00,-26505.80,73494.20,1,0101,证券买入,-26505.80,73494.20,人民币,5.80,4.20,1.60,0.00,0.00,66778001\n"
"2,2024-06-07,,,\t,,,,5120.00,78614.20,2,2041,银行转存,5120.00,78614.20,人民币,0,0,0,0,0,\n"
"3,2024-06-10,000858,五粮液,证券卖出,135.00,200,0,27000.00,26975.30,105589.50,3,0102,证券卖出,26975.30,105589.50,人民币,24.70,4.20,1.60,18.90,0.00,66779001\n"
)
r2 = parse_statement(hs.encode("utf-8"), "hsi.csv")
check("hs: 2 笔成交", len(r2.trades) == 2, f"got {len(r2.trades)}")
check("hs: 费用合计不重复累加", abs(r2.trades[1].fee - 24.70) < 1e-9, f"fee={r2.trades[1].fee}")
check("hs: 深市后缀", r2.trades[0].ts_code == "000858.SZ", r2.trades[0].ts_code)
check("hs: 日期 YYYY-MM-DD", r2.trades[0].trade_date.isoformat() == "2024-06-07")
# ---------- 3) HTML 伪 .xls同花顺导出常见真身 ----------
html = """<html><head><meta charset="gbk"></head><body>
<table>
<tr><td>客户姓名</td><td>测试</td></tr>
<tr><td>成交日期</td><td>业务名称</td><td>证券代码</td><td>证券名称</td><td>成交价格</td><td>成交数量</td><td>成交金额</td><td>手续费</td></tr>
<tr><td>2024/03/15</td><td>证券买入</td><td>300750</td><td>宁德时代</td><td>182.30</td><td>300</td><td>54,690.00</td><td>16.41</td></tr>
<tr><td>2024/03/18</td><td>证券卖出</td><td>300750</td><td>宁德时代</td><td>185.00</td><td>300</td><td>55,500.00</td><td>5.55</td></tr>
</table></body></html>"""
r3 = parse_statement(html.encode("gbk"), "jiaogedan.xls")
check("html: 2 笔成交", len(r3.trades) == 2, f"got {len(r3.trades)}")
check("html: 千分位金额", r3.trades[0].amount == 54690.0, f"{r3.trades[0].amount}")
check("html: 创业板后缀", r3.trades[0].ts_code == "300750.SZ", r3.trades[0].ts_code)
check("html: 斜杠日期", r3.trades[1].trade_date.isoformat() == "2024-03-18")
# ---------- 4) 无业务名称列:发生金额正负判方向(招商式) ----------
zh = (
"证券名称,成交日期,成交价格,成交数量,发生金额,资金余额,合同编号\n"
"贵州茅台,20240102,1680.00,100,-168005.00,200000.00,SZ1000001\n"
"贵州茅台,20240103,1700.50,100,170049.50,370049.50,SZ1000002\n"
)
r4 = parse_statement(zh.encode("utf-8"), "zszs.csv")
check("sign: 2 笔成交", len(r4.trades) == 2, f"got {len(r4.trades)}")
check("sign: 负金额=买入", (r4.trades[0].direction, r4.trades[1].direction) == ("buy", "sell"),
f"{r4.trades[0].direction}/{r4.trades[1].direction}")
# ---------- 5) xlsxopenpyxl 内存构造) ----------
import io # noqa: E402
from openpyxl import Workbook # noqa: E402
wb = Workbook()
ws = wb.active
ws.append(["对账单", None, None])
ws.append(["成交日期", "业务名称", "证券代码", "证券名称", "成交均价", "成交股数", "成交金额", "佣金", "过户费"])
from datetime import datetime as dt # noqa: E402
ws.append([dt(2024, 2, 28, 14, 35, 0), "证券买入", "688981", "中芯国际", 52.80, 200, 10560.00, 2.50, 1.06])
ws.append([dt(2024, 3, 1, 9, 31, 0), "证券卖出", "688981", "中芯国际", 54.10, 200, 10820.00, 2.50, 1.06])
buf = io.BytesIO()
wb.save(buf)
r5 = parse_statement(buf.getvalue(), "sm.xlsx")
check("xlsx: 2 笔成交", len(r5.trades) == 2, f"got {len(r5.trades)}")
check("xlsx: datetime 日期", r5.trades[0].trade_date.isoformat() == "2024-02-28")
check("xlsx: 科创板后缀", r5.trades[0].ts_code == "688981.SH", r5.trades[0].ts_code)
check("xlsx: 佣金+过户费", abs(r5.trades[0].fee - 3.56) < 1e-9, f"fee={r5.trades[0].fee}")
# ---------- 6) 错误分支 ----------
from fastapi import HTTPException # noqa: E402
try:
parse_statement("随便一串不是交割单的文字,1,2,3".encode("utf-8"), "x.csv")
check("garbage: 应 422", False)
except HTTPException as e:
check("garbage: 422", e.status_code == 422)
print()
if FAIL:
print(f"FAIL {len(FAIL)}: {FAIL}")
sys.exit(1)
print("PASS: 交割单解析器全部用例通过")

45
backend/uv.lock generated
View File

@@ -339,6 +339,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/d1/d6/3965ed04c63042e047cb6a3e6ed1a63a35087b6a609aa3a15ed8ac56c221/colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6", size = 25335, upload-time = "2022-10-25T02:36:20.889Z" },
]
[[package]]
name = "et-xmlfile"
version = "2.0.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/d3/38/af70d7ab1ae9d4da450eeec1fa3918940a5fafb9055e934af8d6eb0c2313/et_xmlfile-2.0.0.tar.gz", hash = "sha256:dab3f4764309081ce75662649be815c4c9081e88f0837825f90fd28317d4da54", size = 17234, upload-time = "2024-10-25T17:25:40.039Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c1/8b/5fe2cc11fee489817272089c4203e679c63b570a5aaeb18d852ae3cbba6a/et_xmlfile-2.0.0-py3-none-any.whl", hash = "sha256:7a91720bc756843502c3b7504c77b8fe44217c85c537d85037f0f536151b2caa", size = 18059, upload-time = "2024-10-25T17:25:39.051Z" },
]
[[package]]
name = "fastapi"
version = "0.141.1"
@@ -698,6 +707,18 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/a1/5a/4d2b1601df3602dba7a14f3348ba9bfe94a18adb428e693df6154c293831/numpy-2.5.1-cp314-cp314t-win_arm64.whl", hash = "sha256:5a6db61f9aaa57e369905c67d852045d3c4f7126405b29d09b19dec118e9c9cb", size = 10697674, upload-time = "2026-07-04T17:07:58.506Z" },
]
[[package]]
name = "openpyxl"
version = "3.1.5"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "et-xmlfile" },
]
sdist = { url = "https://files.pythonhosted.org/packages/3d/f9/88d94a75de065ea32619465d2f77b29a0469500e99012523b91cc4141cd1/openpyxl-3.1.5.tar.gz", hash = "sha256:cf0e3cf56142039133628b5acffe8ef0c12bc902d2aadd3e0fe5878dc08d1050", size = 186464, upload-time = "2024-06-28T14:03:44.161Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c0/da/977ded879c29cbd04de313843e76868e6e13408a94ed6b987245dc7c8506/openpyxl-3.1.5-py2.py3-none-any.whl", hash = "sha256:5282c12b107bffeef825f4617dc029afaf41d0ea60823bbb665ef3079dc79de2", size = 250910, upload-time = "2024-06-28T14:03:41.161Z" },
]
[[package]]
name = "pandas"
version = "3.0.5"
@@ -878,6 +899,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/0b/d7/1959b9648791274998a9c3526f6d0ec8fd2233e4d4acce81bbae76b44b2a/python_dotenv-1.2.2-py3-none-any.whl", hash = "sha256:1d8214789a24de455a8b8bd8ae6fe3c6b69a5e3d64aa8a8e5d68e694bbcb285a", size = 22101, upload-time = "2026-03-01T16:00:25.09Z" },
]
[[package]]
name = "python-multipart"
version = "0.0.32"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/5b/42/55c32bb9b12693c092ad250a0e82edb5b31ddeda6eb772de5f308b3804ad/python_multipart-0.0.32.tar.gz", hash = "sha256:be54b7f3fa167bb83e4fcd936b887b708f4e57fe75911c02aebf53efaf8d938e", size = 46881, upload-time = "2026-06-04T16:18:58.647Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/e1/04/e8135ebd1ad02c56ec633277529b2602ff99ff634be76cdba5744cf554fd/python_multipart-0.0.32-py3-none-any.whl", hash = "sha256:ff6d3f776f16878c894e52e107296ffc890e913c611b1a4ec6c44e2821fe2e23", size = 30042, upload-time = "2026-06-04T16:18:57.319Z" },
]
[[package]]
name = "pyyaml"
version = "6.0.3"
@@ -924,6 +954,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" },
]
[[package]]
name = "redis"
version = "8.1.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/a8/99/604f0b666d4c616d891cf77ebb9db6bb21601344c051aebf1b72b9ff915f/redis-8.1.0.tar.gz", hash = "sha256:6e1a19beef9225c83efd689c7e6b7da2d5215b1f42cd13b7fc3714d0a09c7b25", size = 5254356, upload-time = "2026-07-30T08:51:00.269Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/66/9d/c5731f6e3608663d4d3656fd8d3aecee8b509c3082818f5a13eae925baea/redis-8.1.0-py3-none-any.whl", hash = "sha256:a4fe1aac3d3b3cc791d4b3d5931c5a956045dc951ee74d1c913ee3ac4d2ee9fb", size = 560618, upload-time = "2026-07-30T08:50:58.497Z" },
]
[[package]]
name = "requests"
version = "2.34.2"
@@ -1075,9 +1114,12 @@ dependencies = [
{ name = "fastapi" },
{ name = "httpx" },
{ name = "numpy" },
{ name = "openpyxl" },
{ name = "pandas" },
{ name = "pydantic" },
{ name = "pydantic-settings" },
{ name = "python-multipart" },
{ name = "redis" },
{ name = "sqlalchemy" },
{ name = "tushare" },
{ name = "uvicorn", extra = ["standard"] },
@@ -1091,9 +1133,12 @@ requires-dist = [
{ name = "fastapi", specifier = ">=0.115" },
{ name = "httpx", specifier = ">=0.28.1" },
{ name = "numpy", specifier = ">=1.26" },
{ name = "openpyxl", specifier = ">=3.1.5" },
{ name = "pandas", specifier = ">=2.2" },
{ name = "pydantic", specifier = ">=2.7" },
{ name = "pydantic-settings", specifier = ">=2.3" },
{ name = "python-multipart", specifier = ">=0.0.32" },
{ name = "redis", specifier = ">=8.1.0" },
{ name = "sqlalchemy", specifier = ">=2.0" },
{ name = "tushare", specifier = ">=1.4" },
{ name = "uvicorn", extras = ["standard"], specifier = ">=0.30" },