看股功能更新
This commit is contained in:
@@ -3,6 +3,9 @@ TUSHARE_TOKEN=22edda0afe44c0609a187ff1ac0bb2a8fc61430f490ec19f7fec8390
|
||||
DATA_ADJUST=qfq
|
||||
DATA_DEFAULT_START=20200101
|
||||
|
||||
# ---- Redis 读缓存(股票列表/筛选项;留空则不缓存直查数据库)----
|
||||
REDIS_URL=redis://default:26d5c71d57344f37b8b4ddb567f2652f0c7ef41c774284ad@cirry.cn:6379
|
||||
|
||||
# ---- LLM(智能选股;智谱 GLM,OpenAI 兼容协议)----
|
||||
# key 在 https://bigmodel.cn 控制台获取,格式形如 xxxxxxxx.yyyyyyyy(id.secret)
|
||||
LLM_BASE_URL=https://open.bigmodel.cn/api/paas/v4
|
||||
|
||||
@@ -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 # 行业/地域筛选项缓存秒数
|
||||
|
||||
36
backend/alembic/versions/20260815_01_add_adj_factor.py
Normal file
36
backend/alembic/versions/20260815_01_add_adj_factor.py
Normal 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")
|
||||
69
backend/alembic/versions/20260815_02_user_data_tables.py
Normal file
69
backend/alembic/versions/20260815_02_user_data_tables.py
Normal 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")
|
||||
@@ -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_rate(2000 年起)
|
||||
均为可空列——历史回补前为 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")
|
||||
51
backend/alembic/versions/20260815_04_user_trades.py
Normal file
51
backend/alembic/versions/20260815_04_user_trades.py
Normal 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")
|
||||
@@ -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:])
|
||||
@@ -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]
|
||||
|
||||
122
backend/scripts/backfill_adj_factor.py
Normal file
122
backend/scripts/backfill_adj_factor.py
Normal 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))
|
||||
162
backend/scripts/backfill_turnover.py
Normal file
162
backend/scripts/backfill_turnover.py
Normal file
@@ -0,0 +1,162 @@
|
||||
"""全量回补换手率(candles.turnover,单位 %)。
|
||||
|
||||
用法(在 backend 目录下):
|
||||
uv run python scripts/backfill_turnover.py # 从 2000-01-01(daily_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))
|
||||
160
backend/scripts/import_tdx_day.py
Normal file
160
backend/scripts/import_tdx_day.py
Normal 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="不清空任何行,纯 upsert(amount 回补模式)")
|
||||
a = ap.parse_args()
|
||||
asyncio.run(main(Path(a.root), a.symbols or None, a.no_clear))
|
||||
113
backend/scripts/test_trades_parser.py
Normal file
113
backend/scripts/test_trades_parser.py
Normal 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) xlsx(openpyxl 内存构造) ----------
|
||||
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
45
backend/uv.lock
generated
@@ -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" },
|
||||
|
||||
Reference in New Issue
Block a user