diff --git a/backend/.env b/backend/.env index d6aeddc..1011e36 100644 --- a/backend/.env +++ b/backend/.env @@ -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 diff --git a/backend/.env.example b/backend/.env.example index c2d017b..fc5e726 100644 --- a/backend/.env.example +++ b/backend/.env.example @@ -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 # 行业/地域筛选项缓存秒数 diff --git a/backend/alembic/versions/20260815_01_add_adj_factor.py b/backend/alembic/versions/20260815_01_add_adj_factor.py new file mode 100644 index 0000000..b3ba094 --- /dev/null +++ b/backend/alembic/versions/20260815_01_add_adj_factor.py @@ -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") diff --git a/backend/alembic/versions/20260815_02_user_data_tables.py b/backend/alembic/versions/20260815_02_user_data_tables.py new file mode 100644 index 0000000..25d3e4d --- /dev/null +++ b/backend/alembic/versions/20260815_02_user_data_tables.py @@ -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") diff --git a/backend/alembic/versions/20260815_03_candle_amount_turnover.py b/backend/alembic/versions/20260815_03_candle_amount_turnover.py new file mode 100644 index 0000000..186149b --- /dev/null +++ b/backend/alembic/versions/20260815_03_candle_amount_turnover.py @@ -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") diff --git a/backend/alembic/versions/20260815_04_user_trades.py b/backend/alembic/versions/20260815_04_user_trades.py new file mode 100644 index 0000000..63222dd --- /dev/null +++ b/backend/alembic/versions/20260815_04_user_trades.py @@ -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") diff --git a/backend/app/api.py b/backend/app/api.py index 9fa7f7b..c523bd3 100644 --- a/backend/app/api.py +++ b/backend/app/api.py @@ -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) diff --git a/backend/app/backtest/events.py b/backend/app/backtest/events.py new file mode 100644 index 0000000..03073fa --- /dev/null +++ b/backend/app/backtest/events.py @@ -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 diff --git a/backend/app/cache.py b/backend/app/cache.py new file mode 100644 index 0000000..185400c --- /dev/null +++ b/backend/app/cache.py @@ -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 diff --git a/backend/app/config.py b/backend/app/config.py index 7e8111b..fcc04df 100644 --- a/backend/app/config.py +++ b/backend/app/config.py @@ -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 = "" # 留空则智能选股不可用(其余功能不受影响) diff --git a/backend/app/main.py b/backend/app/main.py index ddfd06d..e10866a 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -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( diff --git a/backend/app/models.py b/backend/app/models.py index 80bafe0..4bb9179 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -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" diff --git a/backend/app/schemas.py b/backend/app/schemas.py index ae48d29..b241bf7 100644 --- a/backend/app/schemas.py +++ b/backend/app/schemas.py @@ -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 diff --git a/backend/app/trades.py b/backend/app/trades.py new file mode 100644 index 0000000..f15a169 --- /dev/null +++ b/backend/app/trades.py @@ -0,0 +1,329 @@ +"""交割单解析(券商导出的成交流水 → 结构化买卖记录)。 + +支持三类导出物(按内容嗅探,不信任扩展名): + - CSV/制表符文本(utf-8-sig / gbk / gb18030 自动探测) + - Excel .xlsx(openpyxl;很多券商导出的 .xls 实为 xlsx 或 HTML,先按魔数分流) + - HTML 表格(.xls 常见真身:
| )
+
+列名模糊匹配兼容通达信/恒生/同花顺系的命名差异;业务名称含「买入/卖出」
+才入库,银行转账、配号、利息、红利等非交易行跳过并计数。
+"""
+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" | ||||||||||||||||||||||||||
切。"""
+ 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:])
diff --git a/backend/pyproject.toml b/backend/pyproject.toml
index 25341fb..3f95c7d 100644
--- a/backend/pyproject.toml
+++ b/backend/pyproject.toml
@@ -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]
diff --git a/backend/scripts/backfill_adj_factor.py b/backend/scripts/backfill_adj_factor.py
new file mode 100644
index 0000000..1bdf9b4
--- /dev/null
+++ b/backend/scripts/backfill_adj_factor.py
@@ -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))
diff --git a/backend/scripts/backfill_turnover.py b/backend/scripts/backfill_turnover.py
new file mode 100644
index 0000000..97c57e1
--- /dev/null
+++ b/backend/scripts/backfill_turnover.py
@@ -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))
diff --git a/backend/scripts/import_tdx_day.py b/backend/scripts/import_tdx_day.py
new file mode 100644
index 0000000..80481cc
--- /dev/null
+++ b/backend/scripts/import_tdx_day.py
@@ -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("
+
-
+
-
-
@@ -578,8 +794,33 @@ watch(() => props.subHeights, () => {
+
+
+
+
+
+
+ {{ tradeTip.date }}
+ {{ tradeTip.kind }}
+
+
+
+
+ {{ r.label }}
+ {{ r.text }}
+
+ |