diff --git a/frontend/src/api/paper.ts b/frontend/src/api/paper.ts index 07ef7e5..13fb091 100644 --- a/frontend/src/api/paper.ts +++ b/frontend/src/api/paper.ts @@ -21,6 +21,8 @@ export interface PaperCreate { pool?: string max_pool?: number benchmark?: string + // 撮合引擎(影子柜台 P1):eod_replay=日终回放 / shadow=影子柜台 + engine?: string } export interface PaperAccount { @@ -30,6 +32,7 @@ export interface PaperAccount { interval: string status: string strategy_type?: string + engine?: string symbols?: string initial_capital?: number start_date?: string diff --git a/frontend/src/api/portfolio.ts b/frontend/src/api/portfolio.ts index 8638bfd..213734e 100644 --- a/frontend/src/api/portfolio.ts +++ b/frontend/src/api/portfolio.ts @@ -13,6 +13,8 @@ export interface PortfolioBacktestReq { stamp_duty_rate?: number min_commission?: number slippage?: number + // K线周期(组合回放暂仅日线 d) + interval?: string } export interface EquityPoint { diff --git a/frontend/src/styles/chips.css b/frontend/src/styles/chips.css index 2a259a8..859338d 100644 --- a/frontend/src/styles/chips.css +++ b/frontend/src/styles/chips.css @@ -34,6 +34,12 @@ background: rgba(255, 176, 0, 0.1); border-color: rgba(255, 176, 0, 0.3); } +.chip-shadow { + color: var(--brand); + background: rgba(0, 229, 255, 0.1); + border-color: rgba(0, 229, 255, 0.3); + margin-left: 4px; +} .chip-portfolio { color: var(--amber); background: rgba(255, 176, 0, 0.08); diff --git a/frontend/src/views/paper/List.vue b/frontend/src/views/paper/List.vue index 2897a27..0796593 100644 --- a/frontend/src/views/paper/List.vue +++ b/frontend/src/views/paper/List.vue @@ -188,7 +188,12 @@ async function saveEdit(): Promise { - + diff --git a/frontend/src/views/paper/New.vue b/frontend/src/views/paper/New.vue index fd96ddb..8863254 100644 --- a/frontend/src/views/paper/New.vue +++ b/frontend/src/views/paper/New.vue @@ -115,6 +115,8 @@ onMounted(async () => { const isPortfolio = computed(() => strategyType.value === 'portfolio') const portfolioStrategy = ref('all_weather') +// 撮合引擎:日终回放(NAS 每晚 20:30) / 影子柜台(VPS 盘中实时本地撮合) +const engine = ref<'eod_replay' | 'shadow'>('eod_replay') function setMode(m: string): void { // 组合类型锁定实走:回放卡片不可选(历史回放走「组合回测」页) @@ -149,9 +151,12 @@ async function onSubmit(): Promise { payload.pool = poolForm.pool payload.max_pool = Number(poolForm.max_pool) payload.benchmark = poolForm.benchmark + payload.engine = engine.value } const aid = await createPaper(payload) - ElMessage.success(`已创建模拟盘 #${aid}(今晚 20:30 起每日结算)`) + ElMessage.success(strategyType.value === 'portfolio' && engine.value === 'shadow' + ? `已创建影子柜台模拟盘 #${aid}(VPS 柜台运行期间盘中实时结算)` + : `已创建模拟盘 #${aid}(今晚 20:30 起每日结算)`) router.push(payload.mode === 'live' ? `/paper/live/${aid}` : `/paper/result/${aid}`) } catch (e: unknown) { ElMessage.error(e instanceof Error ? e.message : '创建失败') @@ -245,6 +250,17 @@ function onSymbols(v: string): void { + + + 日终回放 + 影子柜台 + + + {{ engine === 'shadow' + ? 'VPS 盘中实时行情本地撮合(需 VPS 影子柜台进程运行中)' + : '每晚 20:30 全量重放结算(NAS)' }} + + diff --git a/sanguo_api/routes_paper.py b/sanguo_api/routes_paper.py index 79c40ea..b0ebe83 100644 --- a/sanguo_api/routes_paper.py +++ b/sanguo_api/routes_paper.py @@ -56,6 +56,8 @@ class PaperCreateRequest(BaseModel): # 组合策略实走(E1):strategy_type=portfolio 时 mode 必须 live, # strategies[0].name=组合策略名,pool/max_pool/benchmark 进 params strategy_type: str = "cta" + # 撮合引擎(影子柜台 P1):eod_replay=日终回放(NAS 20:30) / shadow=影子柜台(VPS 盘中实时) + engine: str = "eod_replay" pool: str = "hs300_subset" max_pool: int = 30 benchmark: str = "000300.XSHG" @@ -71,6 +73,8 @@ def create_paper(req: PaperCreateRequest): if req.strategy_type == "portfolio": if req.mode != "live": raise HTTPException(400, "组合策略模拟盘仅支持实走(live)模式;历史回放请用「组合回测」") + if req.engine not in ("eod_replay", "shadow"): + raise HTTPException(400, "engine 须为 eod_replay(日终回放) 或 shadow(影子柜台)") payload = req.model_dump() payload["symbols"] = [req.pool] payload["strategies"] = [{ diff --git a/sanguo_trader/persistence.py b/sanguo_trader/persistence.py index d583829..406e7cf 100644 --- a/sanguo_trader/persistence.py +++ b/sanguo_trader/persistence.py @@ -74,6 +74,11 @@ def init_db(db_path: str) -> None: conn.execute("ALTER TABLE paper_accounts ADD COLUMN strategy_type TEXT DEFAULT 'cta'") except sqlite3.OperationalError: pass # 列已存在 + # 迁移:老库补 engine 列(影子柜台 P1:eod_replay=日终回放 / shadow=影子柜台) + try: + conn.execute("ALTER TABLE paper_accounts ADD COLUMN engine TEXT DEFAULT 'eod_replay'") + except sqlite3.OperationalError: + pass # 列已存在 conn.execute("PRAGMA journal_mode=WAL") conn.commit() @@ -85,8 +90,8 @@ def save_account(db_path: str, account: dict[str, Any]) -> int: (task_id, owner_id, name, strategy_type, mode, interval, symbols, strategies, initial_capital, rate, slippage, size, pricetick, stamp_duty_rate, transfer_fee_rate, min_commission, - status, start_date, end_date, created_at, updated_at) - VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)""", + status, start_date, end_date, engine, created_at, updated_at) + VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)""", ( account.get("task_id"), account.get("owner_id", "admin"), account.get("name"), account.get("strategy_type", "cta"), @@ -102,6 +107,7 @@ def save_account(db_path: str, account: dict[str, Any]) -> int: account.get("status", "pending"), account.get("start_date") or account.get("start"), account.get("end_date") or account.get("end"), + account.get("engine", "eod_replay"), _now(), _now(), ), ) diff --git a/sanguo_trader/shadow/__init__.py b/sanguo_trader/shadow/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/sanguo_trader/shadow/__main__.py b/sanguo_trader/shadow/__main__.py new file mode 100644 index 0000000..3523f47 --- /dev/null +++ b/sanguo_trader/shadow/__main__.py @@ -0,0 +1,58 @@ +"""影子柜台 CLI 入口(单实例文件锁防双开重复撮合)。 + +用法: + python -m sanguo_trader.shadow # 读 env(见 runner.py docstring) +""" +from __future__ import annotations + +import logging +import os +import sys +from pathlib import Path + +LOCK_FILE = Path( + os.environ.get("SANGUO_SHADOW_LOCK") + or Path.home() / ".sanguo_shadow_desk.lock" +) + + +def _acquire_lock() -> "object | None": + """单实例锁(Windows/msvcrt 与 POSIX/fcntl 双兼容)。""" + LOCK_FILE.parent.mkdir(parents=True, exist_ok=True) + try: + import fcntl # POSIX + + fh = open(LOCK_FILE, "w") + fcntl.flock(fh, fcntl.LOCK_EX | fcntl.LOCK_NB) + return fh + except ImportError: + pass + except OSError: + return None # 已有实例在跑 + try: + import msvcrt # Windows + + fh = open(LOCK_FILE, "w") + msvcrt.locking(fh.fileno(), msvcrt.LK_NBLCK, 1) + return fh + except (ImportError, OSError): + return None + + +def main() -> int: + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(levelname)s %(name)s: %(message)s", + ) + lock = _acquire_lock() + if lock is None: + print("影子柜台已在运行(锁占用),本次启动退出。", flush=True) + return 0 + from .runner import run_shadow + + run_shadow() + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/sanguo_trader/shadow/broker.py b/sanguo_trader/shadow/broker.py new file mode 100644 index 0000000..def338f --- /dev/null +++ b/sanguo_trader/shadow/broker.py @@ -0,0 +1,228 @@ +"""影子柜台本地模拟 broker(P1,docs/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) diff --git a/sanguo_trader/shadow/runner.py b/sanguo_trader/shadow/runner.py new file mode 100644 index 0000000..d06467d --- /dev/null +++ b/sanguo_trader/shadow/runner.py @@ -0,0 +1,165 @@ +"""影子柜台常驻进程入口(P1,VPS Windows / miniQMT 行情)。 + +与组合实盘(``sanguo_portfolio.runner_live``)同一个 bullet_trade LiveEngine, +唯一区别:broker_factory 换成 ShadowBroker(本地撮合,订单不出门)。 +策略/行情/调度完全同款 → 双轨一致性验证(设计 §8)的基础。 + +环境变量(复用 live_strategy.py 的 SANGUO_LIVE_* 命名 + 影子专属 SANGUO_SHADOW_*): + SANGUO_LIVE_STRATEGY/_MAX_POOL/_BENCHMARK/_CASH 策略配置(live_strategy.py 读) + SANGUO_SHADOW_DB / SANGUO_SHADOW_ACCOUNT_ID 落库目标(paper 库) + SANGUO_SHADOW_COMMISSION/_STAMP/_MIN_COMM/_SLIPPAGE 费率滑点(对齐实盘券商参数) + +手动用法(VPS 交易日): + set SANGUO_LIVE_STRATEGY=all_weather + set SANGUO_SHADOW_DB=C:\\sanguo_vnpy_v2\\data\\paper.db + python -m sanguo_trader.shadow + +不做多账户轮询:MVP 一进程一账户(与 runner_live 一致),多账户由 supervisor +按 paper_accounts(engine='shadow')逐行拉子进程(后续接入)。 +""" +from __future__ import annotations + +# ENV GUARD 必须早于任何 bullet_trade import +import os +os.environ.setdefault("DEFAULT_DATA_PROVIDER", "miniqmt") + +import logging +import threading +import time +from pathlib import Path +from typing import Any, Dict, Optional + +logger = logging.getLogger(__name__) + +# 策略适配文件与组合实盘共用(读 SANGUO_LIVE_* env) +ADAPTER_FILE = Path(__file__).resolve().parents[2] / "sanguo_portfolio" / "live_strategy.py" + + +def shadow_env() -> Dict[str, str]: + """解析影子柜台 env(独立出来便于单测)。""" + return { + "db": os.environ.get("SANGUO_SHADOW_DB", ""), + "account_id": os.environ.get("SANGUO_SHADOW_ACCOUNT_ID", ""), + "commission": os.environ.get("SANGUO_SHADOW_COMMISSION", "0.0003"), + "stamp": os.environ.get("SANGUO_SHADOW_STAMP", "0.001"), + "min_comm": os.environ.get("SANGUO_SHADOW_MIN_COMM", "5"), + "slippage": os.environ.get("SANGUO_SHADOW_SLIPPAGE", "0.001"), + "snapshot_sec": os.environ.get("SANGUO_SHADOW_SNAPSHOT_SEC", "30"), + } + + +def build_price_getter(provider: Any) -> Any: + """从数据 provider 取标的最新价(实时/最新收盘)。返回闭包给 ShadowBroker。""" + + def get_price(security: str) -> Optional[float]: + from datetime import datetime, timedelta + + end = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + start = (datetime.now() - timedelta(days=10)).strftime("%Y-%m-%d") + try: + df = provider.get_price( + security=security, start_date=start, end_date=end, + frequency="daily", fields=["close"], fq="pre", + ) + if df is None or len(df) == 0: + return None + return float(df["close"].iloc[-1]) + except Exception: # noqa: BLE001 - provider 接口差异兜底 + cols = [c for c in ("close", "Close") if c in (df.columns if df is not None else [])] + if cols: + return float(df[cols[0]].iloc[-1]) + return None + + return get_price + + +def _paper_on_trade(db: str, account_id: int, strategy_id: str): + """成交回调:落 paper_trades(与组合实走 EOD 同表,前端模拟盘页直接可见)。""" + from sanguo_trader.persistence import save_trade + + def hook(trade: Dict[str, Any]) -> None: + side = trade["side"] + save_trade(db, account_id, { + "strategy_id": strategy_id, + "datetime": trade["datetime"], + "symbol": trade["security"], + "direction": "long" if side == "buy" else "short", + "offset": "open" if side == "buy" else "close", + "match_session": "shadow_realtime", + "price": trade["price"], + "volume": trade["amount"], + "commission": trade["commission"], + "stamp_duty": trade["stamp_duty"], + "bar_date": trade["datetime"][:10], + }) + + return hook + + +def _snapshot_loop(broker: Any, db: str, account_id: int, + interval_sec: float = 30.0) -> None: + """后台线程:定期把影子账户快照落 paper_positions/paper_daily_balance。""" + from sanguo_trader.persistence import save_daily_balance, save_positions + + while True: + time.sleep(interval_sec) + try: + info = broker.get_account_info() + positions = { + p["security"]: {"volume": float(p["amount"]), "frozen": 0.0, + "avg_price": p["avg_cost"]} + for p in info["positions"] + } + save_positions(db, account_id, "account", positions, + date=broker.trades[-1]["datetime"][:10] if broker.trades else "") + save_daily_balance( + db, account_id, info.get("as_of", ""), + cash=info["available_cash"], market_value=info["market_value"], + total_equity=info["total_value"], + ) + except Exception as exc: # noqa: BLE001 - 落库失败不中断柜台 + logger.warning("[shadow-snapshot] 落库失败 (account=%s): %s", account_id, exc) + + +def run_shadow(provider_config: Optional[Dict[str, Any]] = None) -> None: + """装配 LiveEngine(影子 broker)并 run(阻塞)。""" + from bullet_trade.core.live_engine import LiveEngine # type: ignore + from bullet_trade.data.api import set_data_provider # type: ignore + + from sanguo_portfolio.runner_live import build_provider + from .broker import ShadowBroker + + cfg = shadow_env() + from sanguo_portfolio.runner_live import live_env + le = live_env() + + provider = build_provider(provider_config) + set_data_provider(provider) + + broker = ShadowBroker( + initial_cash=float(le["cash"]), + commission_rate=float(cfg["commission"]), + stamp_duty_rate=float(cfg["stamp"]), + min_commission=float(cfg["min_comm"]), + slippage=float(cfg["slippage"]), + price_getter=build_price_getter(provider), + on_trade=_paper_on_trade(cfg["db"], int(cfg["account_id"]), le["strategy"]) + if cfg["db"] and cfg["account_id"] else None, + ) + logger.info( + "影子柜台启动: strategy=%s cash=%s 费率=佣金%s/印花%s/最低%s 滑点%s db=%s", + le["strategy"], le["cash"], cfg["commission"], cfg["stamp"], + cfg["min_comm"], cfg["slippage"], cfg["db"] or "(不落库)", + ) + + engine = LiveEngine(ADAPTER_FILE, broker_factory=lambda: broker) + + if cfg["db"] and cfg["account_id"]: + t = threading.Thread( + target=_snapshot_loop, + args=(broker, cfg["db"], int(cfg["account_id"]), float(cfg["snapshot_sec"])), + daemon=True, name="shadow-snapshot", + ) + t.start() + + engine.run() diff --git a/tests/trader/test_shadow_broker.py b/tests/trader/test_shadow_broker.py new file mode 100644 index 0000000..207f1fa --- /dev/null +++ b/tests/trader/test_shadow_broker.py @@ -0,0 +1,135 @@ +"""ShadowBroker(影子柜台本地撮合)单元测试。 + +纯逻辑测试:不依赖 bullet_trade/xtquant,价格由固定 price_getter 注入。 +""" +from __future__ import annotations + +import asyncio +from datetime import datetime + +import pytest + +from sanguo_trader.shadow.broker import ShadowBroker + + +def _mk_broker(cash: float = 100_000.0, **kw) -> ShadowBroker: + prices = kw.pop("prices", {"600519.SH": 100.0}) + fixed_now = kw.pop("now", datetime(2026, 8, 14, 10, 0, 0)) + return ShadowBroker( + initial_cash=cash, + price_getter=lambda s: prices.get(s), + now_provider=lambda: fixed_now, + **kw, + ) + + +def _buy(b: ShadowBroker, sec: str, amt: int, px: float | None = None): + return asyncio.run(b.buy(sec, amt, px)) + + +def _sell(b: ShadowBroker, sec: str, amt: int, px: float | None = None): + return asyncio.run(b.sell(sec, amt, px)) + + +def test_buy_fills_with_commission_and_slippage(): + b = _mk_broker(cash=100_000, slippage=0.001, commission_rate=0.0003, min_commission=5) + oid = _buy(b, "600519.SH", 100) + assert b.orders[oid]["status"] == "filled" + fill = b.orders[oid]["filled_price"] + assert fill == pytest.approx(100.0 * 1.001, abs=0.01) # 买入上浮滑点 + # 现金扣减 = 全额 + 佣金(低于最低佣金取 5 元) + commission = max(100 * fill * 0.0003, 5.0) + assert b.cash == pytest.approx(100_000 - 100 * fill - commission) + assert b.positions["600519.SH"]["amount"] == 100 + assert b.positions["600519.SH"]["avg_cost"] == pytest.approx(fill) + + +def test_sell_charges_stamp_duty_and_slippage_down(): + b = _mk_broker(cash=100_000, slippage=0.001, stamp_duty_rate=0.001) + _buy(b, "600519.SH", 200, px=100.0) # 固定委托价,滑点仍生效 + cash_after_buy = b.cash + # T+1:当日买入不可卖 → 先模拟次日(before_open 清锁) + b.before_open() + oid = _sell(b, "600519.SH", 200, px=100.0) + assert b.orders[oid]["status"] == "filled" + fill = b.orders[oid]["filled_price"] + assert fill == pytest.approx(100.0 * 0.999, abs=0.01) # 卖出下压滑点 + gross = 200 * fill + commission = max(gross * 0.0003, 5.0) + stamp = gross * 0.001 + assert b.cash == pytest.approx(cash_after_buy + gross - commission - stamp) + assert "600519.SH" not in b.positions # 清仓移除 + + +def test_t1_blocks_same_day_sell(): + b = _mk_broker() + _buy(b, "600519.SH", 200, px=100.0) + oid = _sell(b, "600519.SH", 200, px=100.0) # 当日卖 → 拒 + assert b.orders[oid]["status"] == "rejected" + assert "T+1" in b.orders[oid]["reject_reason"] + # 次日可卖 + b.before_open() + oid2 = _sell(b, "600519.SH", 200, px=100.0) + assert b.orders[oid2]["status"] == "filled" + + +def test_insufficient_cash_rejects(): + b = _mk_broker(cash=5_000) + oid = _buy(b, "600519.SH", 100) # 需约 1 万 + assert b.orders[oid]["status"] == "rejected" + assert "资金不足" in b.orders[oid]["reject_reason"] + assert b.cash == 5_000 # 拒单不动账 + + +def test_odd_lot_floors_to_100(): + b = _mk_broker(cash=1_000_000) + oid = _buy(b, "600519.SH", 250) # → 200 + assert b.orders[oid]["status"] == "filled" + assert b.orders[oid]["filled_amount"] == 200 + oid2 = _buy(b, "600519.SH", 50) # 不足一手 → 拒 + assert b.orders[oid2]["status"] == "rejected" + + +def test_no_price_rejects(): + b = _mk_broker(prices={}) + oid = _buy(b, "600519.SH", 100) + assert b.orders[oid]["status"] == "rejected" + assert "无参考价" in b.orders[oid]["reject_reason"] + + +def test_on_trade_callback_receives_fills(): + seen: list[dict] = [] + b = ShadowBroker( + initial_cash=100_000, + price_getter=lambda s: 10.0, + on_trade=seen.append, + ) + _buy(b, "600519.SH", 100) + assert len(seen) == 1 + t = seen[0] + assert t["side"] == "buy" and t["amount"] == 100 and t["price"] == pytest.approx(10.0) + + +def test_get_account_info_totals(): + b = _mk_broker(cash=100_000) + _buy(b, "600519.SH", 100, px=100.0) + info = b.get_account_info() + assert info["available_cash"] == pytest.approx(100_000 - 100 * 100.0 - max(100 * 100 * 0.0003, 5)) + assert info["market_value"] == pytest.approx(100 * 100.0) # price_getter=100 + assert info["total_value"] == pytest.approx(info["available_cash"] + info["market_value"]) + + +def test_avg_cost_weighted_on_second_buy(): + b = _mk_broker(cash=1_000_000, slippage=0.0) + _buy(b, "600519.SH", 100, px=100.0) + _buy(b, "600519.SH", 100, px=110.0) + pos = b.positions["600519.SH"] + assert pos["amount"] == 200 + assert pos["avg_cost"] == pytest.approx(105.0) + + +def test_cancel_always_false_and_open_orders_empty(): + b = _mk_broker() + _buy(b, "600519.SH", 100, px=100.0) + assert asyncio.run(b.cancel_order("whatever")) is False + assert b.get_open_orders() == [] # 即时成交,无挂单