Files
claude_dev b66a6c5958
CI/CD / test (push) Successful in 11s
CI/CD / nas-deploy (push) Successful in 29s
CI/CD / nas-verify (push) Successful in 4m11s
fix(ashare_engine): load_data 直查 dbbardata 绕 vnpy peewee(治 CTA 死算)
根因:vnpy 父类 load_data 走 get_database().load_bar_data()(peewee ORM),
get_database init 时 create_indexes 对 26G dbbardata 库 → spawn 建 26G
索引 CPU 死算 1:27+ / web 并发 database is locked(2026-08-02 根因,
vnpy_db 指向 v2 大库 2cb2ab0 后触发)。load_data 读 0 bars → 回测空/卡。

修复:load_data 统一直查 dbbardata(原仅 5m/15m 走 _load_intraday_data,
扩展到日线 d),绕 vnpy peewee get_database。sqlite3 busy_timeout=30s
防并发锁。d → BarData(interval=Interval.DAILY)。

验证:日线 600519 2024H1 load_data 0.1s bars=116(was database locked
5.2s/0 bars) interval=DAILY;run_backtesting 5.6s(was 死算 1:27+);
000001 compute_metrics 完成 daily_df=(116,25) benchmark=116(5图metrics恢复)。
2026-08-02 21:48:33 +08:00

274 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""A 股适配层:子类化 vnpy_ctastrategy BacktestingEngine + DailyResult。
vnpy 源码零修改,全部覆写在本文件。
- AShareDailyResultA 股费用模型(佣金保底、印花税卖方、过户费沪市)
- AShareBacktestingEngine
- send_order 覆写 → 做空拦截(SHORT+OPEN 拒单)+ 定寸重算 volume
- update_daily_close 覆写 → 工厂换 AShareDailyResult(父类在 :647 实例化 DailyResult
- load_data 覆写 → 支持 5m/15m 周期(vnpy Interval enum 不认 '5m'/'15m'
"""
from __future__ import annotations
import logging
from datetime import datetime
from vnpy_ctastrategy.backtesting import BacktestingEngine, DailyResult
from vnpy.trader.constant import Direction, Offset, Interval, Exchange
logger = logging.getLogger(__name__)
class AShareDailyResult(DailyResult):
"""A 股日度盈亏:佣金双边保底 + 印花税卖方 + 过户费沪市。
父类 calculate_pnl 用 turnover*rate 算佣金(单边商),A 股实际:
- 佣金 = max(turnover * commission_rate, min_commission),买卖都收
- 印花税 = turnover * stamp_duty_rate,卖方 onlydirection==SHORT
- 过户费 = turnover * transfer_fee_rate,沪市 only
"""
def __init__(
self,
date,
close_price: float,
commission_rate: float = 0.00025,
min_commission: float = 5.0,
stamp_duty_rate: float = 0.0005,
transfer_fee_rate: float = 0.00001,
is_sse: bool = False,
) -> None:
super().__init__(date, close_price)
self.commission_rate: float = commission_rate
self.min_commission: float = min_commission
self.stamp_duty_rate: float = stamp_duty_rate
self.transfer_fee_rate: float = transfer_fee_rate
self.is_sse: bool = is_sse
# 费用明细(持久化可见——父类 calculate_result 会遍历 __dict__ 落列)
self.stamp_duty: float = 0.0
self.transfer_fee: float = 0.0
def calculate_pnl(
self,
pre_close: float,
start_pos: float,
size: float,
rate: float,
slippage: float,
) -> None:
"""覆写父类:用 A 股费用模型替换 commission = turnover * rate。
签名与父类一致(calculate_result 传 self.rate/self.slippage/self.size),
但 rate 参数被忽略——佣金由 self.commission_rate + self.min_commission 决定。
"""
# 首日无 pre_close 时用 1 防除零(与父类逻辑一致)
self.pre_close = pre_close if pre_close else 1
self.start_pos = start_pos
self.end_pos = start_pos
self.holding_pnl = self.start_pos * (self.close_price - self.pre_close) * size
self.trade_count = len(self.trades)
for trade in self.trades:
if trade.direction == Direction.LONG:
pos_change = trade.volume
else:
pos_change = -trade.volume
self.end_pos += pos_change
turnover: float = trade.volume * size * trade.price
self.trading_pnl += pos_change * (self.close_price - trade.price) * size
self.slippage += trade.volume * size * slippage
self.turnover += turnover
# A 股佣金:双边,最低 min_commission 元
self.commission += max(turnover * self.commission_rate, self.min_commission)
# 印花税:卖方 onlySHORT = 卖出)
if trade.direction == Direction.SHORT:
self.stamp_duty += turnover * self.stamp_duty_rate
# 过户费:沪市 only
if self.is_sse:
self.transfer_fee += turnover * self.transfer_fee_rate
# net_pnl 扣除全部费用
self.total_pnl = self.trading_pnl + self.holding_pnl
self.net_pnl = (
self.total_pnl
- self.commission
- self.slippage
- self.stamp_duty
- self.transfer_fee
)
class AShareBacktestingEngine(BacktestingEngine):
"""A 股回测引擎:long-only 拦截 + A 股费用。
覆写点:
1. send_order → 拦截 SHORT+OPENA股不可做空)
2. update_daily_close → 工厂换 AShareDailyResult(父类 :647 实例化 DailyResult
定寸(C1)通过 engine.size = N 实现,不在此处处理:
cta_engine 在 load_data 后取首根 bar close 算 N = floor(capital*pct/close/100)*100
设置 engine.size = N。vnpy 的 turnover/PnL 自动 ×size,策略 volume 保持 1 手=N 股=满仓。
"""
def __init__(self) -> None:
super().__init__()
# A 股费用参数(默认值,cta_engine 可覆盖)
self.commission_rate: float = 0.00025 # 万 2.5
self.min_commission: float = 5.0 # 最低 5 元
self.stamp_duty_rate: float = 0.0005 # 卖方 0.05%
self.transfer_fee_rate: float = 0.00001 # 沪市 0.001%
self.is_sse: bool = False
# 5m/15m 适配:vnpy Interval enum 不认 '5m'/'15m'cta_engine 把 engine.interval
# 映射成 Interval.MINUTE 让父类 set_parameters 校验通过,真实 DB interval 字符串
# 存此字段供 load_data 自定义路径使用。默认 "d" 走 vnpy 原生日线路径。
self.raw_interval: str = "d"
self.sqlite_db_path: str | None = None # 由 cta_engine 注入(_dcfg.data_paths["vnpy_db"]
def send_order(
self,
strategy,
direction: Direction,
offset: Offset,
price: float,
volume: float,
stop: bool,
lock: bool,
net: bool,
) -> list:
"""覆写父类 send_order:做空拦截(C2)。
C2 做空拦截:SSE/SZSE 标的不可做空,SHORT+OPEN 直接拒单。
SHORT+CLOSE(平多)允许。
volume 不动——定寸由 engine.size = N 实现(见类文档)。
"""
if direction == Direction.SHORT and offset == Offset.OPEN:
logger.warning(
"A股不支持做空,拒单: direction=%s offset=%s price=%s volume=%s",
direction, offset, price, volume,
)
return []
return super().send_order(
strategy, direction, offset, price, volume, stop, lock, net
)
def update_daily_close(self, price: float) -> None:
"""覆写父类工厂方法:用 AShareDailyResult 替换 DailyResult。
父类原实现(backtesting.py:639-647):
daily_result = self.daily_results.get(d)
if daily_result: daily_result.close_price = price
else: self.daily_results[d] = DailyResult(d, price)
"""
d = self.datetime.date()
daily_result = self.daily_results.get(d, None)
if daily_result:
daily_result.close_price = price
else:
self.daily_results[d] = AShareDailyResult(
d,
price,
commission_rate=self.commission_rate,
min_commission=self.min_commission,
stamp_duty_rate=self.stamp_duty_rate,
transfer_fee_rate=self.transfer_fee_rate,
is_sse=self.is_sse,
)
def load_data(self) -> None:
"""覆写父类:直查 dbbardata(绕 vnpy peewee get_database 对大库锁/死算 + enum 限制)。
vnpy Interval enum 只有 MINUTE('1m')/HOUR('1h')/DAILY('d') 等,不认 '5m'/'15m'
父类 load_data 调 ``INTERVAL_DELTA_MAP[self.interval]`` 和
``load_bar_data(..., self.interval, ...)`` 都依赖 enum,传 '15m' 会 ValueError。
适配思路:self.interval 已被父类 set_parameters 映射成 Interval.MINUTEenum
校验通过),真实 DB interval 存 self.raw_interval。当 raw_interval 是 5m/15m
时,直接 sqlite 查 dbbardata 表,把每行转 BarData(interval=Interval.MINUTE)
绕过 vnpy database_manager 的 enum 限制;其余周期(d/1m 等)走 vnpy 原生路径。
"""
# 统一直查 dbbardata(绕 vnpy peewee get_database 对大库 create_indexes 锁/死算;
# 2026-08-02 根因:vnpy_db 指向 v2 26G 库后 super().load_data 建 26G 索引 spawn 死算 1:27+
self._load_intraday_data()
def _load_intraday_data(self) -> None:
"""5m/15m 直查 SQLite → BarData(MINUTE),绕过 vnpy enum 限制。"""
from vnpy.trader.object import BarData
import sqlite3
self.output(f"开始加载 {self.raw_interval} 历史数据(ashare 适配)")
if not self.end:
self.end = datetime.now()
if self.start >= self.end:
self.output("起始日期必须小于结束日期")
return
db_path = self.sqlite_db_path
if not db_path:
try:
from vnpy.trader.setting import SETTINGS
db_path = SETTINGS.get("database.database")
except Exception:
db_path = None
if not db_path:
raise RuntimeError(
"5m/15m 回测需要 sqlite_db_path(或 SETTINGS['database.database']),"
"cta_engine 应在 load_data 前注入。"
)
symbol, exchange_str = self.vt_symbol.split(".")
# peewee DateTimeField 存的是 ISO 字符串;start/end 用 datetime 比较即可
# (SQLite 会把参数转成可比较的字符串形式)。
conn = sqlite3.connect(db_path, timeout=30)
conn.execute("PRAGMA busy_timeout = 30000")
try:
cur = conn.execute(
"SELECT datetime, volume, turnover, open_interest, "
"open_price, high_price, low_price, close_price "
"FROM dbbardata "
"WHERE symbol=? AND exchange=? AND interval=? "
"AND datetime>=? AND datetime<=? "
"ORDER BY datetime",
(symbol, exchange_str, self.raw_interval, self.start, self.end),
)
rows = cur.fetchall()
finally:
conn.close()
exchange = Exchange(exchange_str)
bars: list[BarData] = []
for dt, vol, turnover, oi, o, h, l, c in rows:
if isinstance(dt, str):
try:
dt = datetime.fromisoformat(dt)
except ValueError:
continue
# vnpy_sqlite save 时 convert_tz 改成 UTC,回测时按本地时间跑即可
# (日线回测也是直接读 DB datetime,行为一致)。
bars.append(BarData(
symbol=symbol,
exchange=exchange,
datetime=dt,
interval=(Interval.DAILY if self.raw_interval == "d" else Interval.MINUTE),
volume=float(vol or 0),
turnover=float(turnover or 0),
open_interest=float(oi or 0),
open_price=float(o or 0),
high_price=float(h or 0),
low_price=float(l or 0),
close_price=float(c or 0),
gateway_name="sqlite",
))
self.history_data.clear()
self.history_data.extend(bars)
self.output(f"历史数据加载完成,数据量:{len(self.history_data)}")