Files
sanguo_vnpy_v2/sanguo_trader/shadow/broker.py
T

229 lines
9.8 KiB
Python
Raw 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.
"""影子柜台本地模拟 brokerP1docs/design/paper-shadow-desk-design.md §3.2)。
挂在 bullet_trade LiveEngine 的 broker_factory 上:策略下单不出门,
由本 broker 以「下单时刻实时价 ± 滑点」本地撮合,A 股费用/整手/T+1 对齐。
与实盘(QmtBroker)同接口(BrokerBase) → 同一个 LiveEngine 两种柜台,
这是「双轨一致性验证」(§8)的基础:同策略同参数分别接真/假 broker 并跑对账。
价格来源由 price_getter 注入(通常=数据 provider 最新收盘/实时价)
成交回调 on_trade 注入(落 paper_trades 表)。
"""
from __future__ import annotations
import logging
import uuid
from datetime import datetime
from typing import Any, Callable, Dict, List, Optional
logger = logging.getLogger(__name__)
LOT = 100 # A 股整手
class ShadowBroker: # noqa: R0903 - 仅实现 BrokerBase 协议(bullet_trade duck-typed)
"""本地虚拟账户撮合台。不继承 BrokerBase(避免硬依赖 bullet_trade import 顺序),
LiveEngine 按 duck-typed 协议调用。"""
def __init__(
self,
initial_cash: float = 1_000_000.0,
*,
commission_rate: float = 0.0003,
stamp_duty_rate: float = 0.001,
min_commission: float = 5.0,
slippage: float = 0.0,
price_getter: Optional[Callable[[str], Optional[float]]] = None,
on_trade: Optional[Callable[[Dict[str, Any]], None]] = None,
now_provider: Optional[Callable[[], datetime]] = None,
) -> None:
self.initial_cash = float(initial_cash)
self.cash = float(initial_cash)
self.commission_rate = float(commission_rate)
self.stamp_duty_rate = float(stamp_duty_rate)
self.min_commission = float(min_commission)
self.slippage = float(slippage)
self.price_getter = price_getter
self.on_trade = on_trade
self._now = now_provider or datetime.now
self._connected = True # 本地柜台永远"在线"
# security -> {"amount": int, "avg_cost": float}
self.positions: Dict[str, Dict[str, Any]] = {}
# T+1:今日买入数量(security -> int),before_open 清零
self._today_bought: Dict[str, int] = {}
self._today: str = ""
self.orders: Dict[str, Dict[str, Any]] = {}
self.trades: List[Dict[str, Any]] = []
# ===== 生命周期 =====
def connect(self) -> bool:
return True
def disconnect(self) -> bool:
return True
def is_connected(self) -> bool:
return True
def heartbeat(self) -> None:
return None
def before_open(self) -> None:
"""每个交易日开盘前:清 T+1 买入记录(昨日买的今天可卖)。"""
self._today_bought = {}
self._today = self._now().strftime("%Y-%m-%d")
def after_close(self) -> None:
return None
# ===== 行情 =====
def _ref_price(self, security: str, price: Optional[float]) -> Optional[float]:
ref = price if price and price > 0 else None
if ref is None and self.price_getter is not None:
try:
ref = self.price_getter(security)
except Exception as exc: # noqa: BLE001 - 行情失败拒单而非崩柜台
logger.warning("[shadow] 取价失败 %s: %s", security, exc)
ref = None
return ref
# ===== 下单(即时全额成交) =====
async def buy(self, security: str, amount: int, price: Optional[float] = None,
wait_timeout: Optional[float] = None, remark: Optional[str] = None,
*, market: bool = False) -> str:
order_id = self._new_order("buy", security, amount, price)
ref = self._ref_price(security, price)
if ref is None or ref <= 0:
return self._reject(order_id, "无参考价")
amount = int(amount)
if amount <= 0:
return self._reject(order_id, "数量非法")
amount = amount - amount % LOT # 整手向下取
if amount <= 0:
return self._reject(order_id, "不足一手(100股)")
fill = round(ref * (1 + self.slippage) + 1e-9, 2) # 买入价上浮滑点
gross = amount * fill
commission = max(gross * self.commission_rate, self.min_commission)
if self.cash < gross + commission:
return self._reject(order_id, f"资金不足 需{gross + commission:.2f}{self.cash:.2f}")
self.cash -= gross + commission
pos = self.positions.setdefault(security, {"amount": 0, "avg_cost": 0.0})
old_amt, old_cost = pos["amount"], pos["avg_cost"]
pos["amount"] = old_amt + amount
pos["avg_cost"] = (old_amt * old_cost + gross) / pos["amount"]
self._today_bought[security] = self._today_bought.get(security, 0) + amount
self._fill(order_id, security, "buy", amount, fill, commission, 0.0)
return order_id
async def sell(self, security: str, amount: int, price: Optional[float] = None,
wait_timeout: Optional[float] = None, remark: Optional[str] = None,
*, market: bool = False) -> str:
order_id = self._new_order("sell", security, amount, price)
ref = self._ref_price(security, price)
if ref is None or ref <= 0:
return self._reject(order_id, "无参考价")
amount = int(amount)
pos = self.positions.get(security)
held = int(pos["amount"]) if pos else 0
if amount <= 0 or held <= 0:
return self._reject(order_id, "无持仓")
# T+1:今日买入部分不可卖
locked = self._today_bought.get(security, 0)
sellable = max(held - locked, 0)
if amount > sellable:
amount = sellable
amount = amount - amount % LOT
if amount <= 0:
return self._reject(order_id, f"可卖不足(T+1锁定{locked}股)")
fill = round(ref * (1 - self.slippage) - 1e-9, 2) # 卖出价下压滑点
gross = amount * fill
commission = max(gross * self.commission_rate, self.min_commission)
stamp_duty = gross * self.stamp_duty_rate
self.cash += gross - commission - stamp_duty
pos["amount"] = held - amount
if pos["amount"] <= 0:
self.positions.pop(security, None)
self._fill(order_id, security, "sell", amount, fill, commission, stamp_duty)
return order_id
async def cancel_order(self, order_id: str) -> bool:
# 即时全额成交,无可撤单
return False
async def get_order_status(self, order_id: str) -> Dict[str, Any]:
st = self.orders.get(order_id) or {"order_id": order_id, "status": "not_found"}
return dict(st)
def get_orders(self, order_id=None, security=None, status=None,
from_broker: bool = False) -> List[Dict[str, Any]]:
rows = [dict(o) for o in self.orders.values()
if (order_id is None or o["order_id"] == order_id)
and (security is None or o["security"] == security)]
return rows
def get_open_orders(self) -> List[Dict[str, Any]]:
return [dict(o) for o in self.orders.values() if o["status"] == "open"]
def get_trades(self, order_id=None, security=None) -> List[Dict[str, Any]]:
return [dict(t) for t in self.trades
if (order_id is None or t["order_id"] == order_id)
and (security is None or t["security"] == security)]
# ===== 账户 =====
def get_positions(self) -> List[Dict[str, Any]]:
out = []
for sym, pos in self.positions.items():
px = self._ref_price(sym, None) or pos["avg_cost"]
out.append({"security": sym, "amount": pos["amount"],
"avg_cost": round(pos["avg_cost"], 6),
"market_value": pos["amount"] * px,
"price": px})
return out
def get_account_info(self) -> Dict[str, Any]:
positions = self.get_positions()
mv = sum(p["market_value"] for p in positions)
return {
"total_value": self.cash + mv,
"available_cash": self.cash,
"positions": positions,
"market_value": mv,
}
# ===== 内部 =====
def _new_order(self, side: str, security: str, amount: int,
price: Optional[float]) -> str:
order_id = f"shadow_{uuid.uuid4().hex[:12]}"
self.orders[order_id] = {
"order_id": order_id, "status": "open", "side": side,
"security": security, "amount": int(amount),
"price": price, "created_at": self._now().isoformat(timespec="seconds"),
}
return order_id
def _reject(self, order_id: str, reason: str) -> str:
o = self.orders[order_id]
o["status"] = "rejected"
o["reject_reason"] = reason
logger.info("[shadow] 拒单 %s %s %s: %s", o["side"], o["security"], o["amount"], reason)
return order_id
def _fill(self, order_id: str, security: str, side: str, amount: int,
fill: float, commission: float, stamp_duty: float) -> None:
o = self.orders[order_id]
o.update(status="filled", filled_amount=amount, filled_price=fill)
trade = {
"order_id": order_id, "security": security, "side": side,
"amount": amount, "price": fill, "commission": round(commission, 2),
"stamp_duty": round(stamp_duty, 2),
"datetime": self._now().strftime("%Y-%m-%d %H:%M:%S"),
}
self.trades.append(trade)
logger.info("[shadow] 成交 %s %s %d股 @%.2f%.2f",
side, security, amount, fill, commission + stamp_duty)
if self.on_trade is not None:
try:
self.on_trade(trade)
except Exception as exc: # noqa: BLE001 - 落库失败不阻断撮合
logger.warning("[shadow] on_trade 回调失败: %s", exc)