"""A 股适配层:子类化 vnpy_ctastrategy BacktestingEngine + DailyResult。 vnpy 源码零修改,全部覆写在本文件。 - AShareDailyResult:A 股费用模型(佣金保底、印花税卖方、过户费沪市) - 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,卖方 only(direction==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) # 印花税:卖方 only(SHORT = 卖出) 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+OPEN(A股不可做空) 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: """覆写父类:支持 5m/15m 周期(绕过 vnpy Interval 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.MINUTE(enum 校验通过),真实 DB interval 存 self.raw_interval。当 raw_interval 是 5m/15m 时,直接 sqlite 查 dbbardata 表,把每行转 BarData(interval=Interval.MINUTE), 绕过 vnpy database_manager 的 enum 限制;其余周期(d/1m 等)走 vnpy 原生路径。 """ if self.raw_interval not in ("5m", "15m"): super().load_data() return 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) 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.MINUTE, # 5m/15m 不在 enum,统一标 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)}")