diff --git a/sanguo_portfolio/strategies/channel_test.py b/sanguo_portfolio/strategies/channel_test.py index 054af84..9fe3a2e 100644 --- a/sanguo_portfolio/strategies/channel_test.py +++ b/sanguo_portfolio/strategies/channel_test.py @@ -1,48 +1,53 @@ """通路测试策略(影子柜台 vs 实盘 双轨验证专用,docs/design/paper-shadow-desk-design.md §8)。 -不是为赚钱,是为**把买卖全通路在真实 miniQMT 数据上跑通**,让 ShadowBroker 与 -QmtBroker 在同一段时间、同一批订单上各自成交,事后对账验证通路正确性。 +不是为赚钱,是为**把买卖全通路在真实 miniQMT 数据上每天跑满一遍**,让 ShadowBroker +与 QmtBroker 在同一段时间、同一批订单上各自成交,事后对账验证通路正确性。 -设计目标(每个调仓日都尽量触发): -- 买:用 order_target_value 等权买入若干只 → 走买入 + 整手(100) + 资金扣减 + 均价 -- 卖:次日先清掉上轮非目标持仓 → 走卖出 + 印花税 + T+1 可卖 -- 轮换:目标集每个周期偏移一格 → 既有卖(旧)又有买(新),长期跑必两端都覆盖 -- T+1 拒单探针:买入后立即试卖同一只(当日) → 两端都应被 T+1 拒,验证拒单通路一致 -- 上涨/下跌/停牌/涨跌停:靠自然行情出现,差异进双轨对账报告(P1.3 涨跌停拦截随影子柜台完善) +universe 按资产类型分组(每类多只,轮换时跨类型取样 → 各类型买卖都覆盖): +- 宽基 ETF / 行业 ETF / 跨境·商品 ETF / 主板蓝筹 / 主板中盘 / 创业板 +- ⚠️ 只含主板+创业板个股(实盘账户无科创/北交权限,铁律);ETF 无板块权限限制 -universe 默认几只高流动性 ETF/蓝筹(miniQMT 必有数据、好成交),可经 env/max_pool 调。 +每日场景覆盖(多个盘中时点触发,15m 线账户建议配 interval=15m): +- 09:35 主调仓:卖掉全部非目标持仓(卖全量) + 等权买入 6 只跨类型标的(买/整手/资金) +- 10:45 部分调仓:对已持仓标的加减仓(部分买卖,非清仓式) +- 13:45 卖后买:清掉 1 只持仓换买另 1 只(当日资金复用;清的是昨日仓,T+1 可卖) +- 14:30 T+1 探针:当日买入立即试卖 → 两端都应被 T+1 拒(验证拒单通路) +- 涨跌停/停牌/部分成交:靠真实行情自然出现,差异进双轨对账报告 """ from __future__ import annotations import logging from dataclasses import dataclass, field -from typing import Any, List, Optional +from typing import Any, Dict, List, Optional from .all_weather import BrokerFacade, _available_cash, _get_positions logger = logging.getLogger(__name__) -# 默认 universe:流动性好的 ETF + 蓝筹,miniQMT 必有数据,实盘也容易成交 -_DEFAULT_UNIVERSE: List[str] = [ - "510300.XSHG", # 沪深300ETF - "510050.XSHG", # 上证50ETF - "159915.XSHE", # 创业板ETF - "510500.XSHG", # 中证500ETF - "588000.XSHG", # 科创50ETF -] +# universe 按类型分组;轮换时从每组轮流取 → 每天的持仓组合跨类型 +UNIVERSE_BY_TYPE: Dict[str, List[str]] = { + "宽基ETF": ["510300.XSHG", "510050.XSHG", "510500.XSHG", "159915.XSHE"], + "行业ETF": ["512880.XSHG", "512690.XSHG", "515790.XSHG"], + "跨境商品ETF": ["513100.XSHG", "513030.XSHG", "518880.XSHG"], + "主板蓝筹": ["600519.XSHG", "601318.XSHG", "600036.XSHG"], + "主板中盘": ["000001.XSHE", "600100.XSHG", "601668.XSHG"], + "创业板": ["300750.XSHE", "300059.XSHE", "002415.XSHE"], # 002=中小板并入主板 +} +_TYPE_ORDER = list(UNIVERSE_BY_TYPE.keys()) @dataclass class ChannelTestConfig: - universe: List[str] = field(default_factory=lambda: list(_DEFAULT_UNIVERSE)) - hold_n: int = 2 # 每轮等权持有几只 - period: int = 1 # 每 N 个交易日轮换一次(1=每日) - probe_t1: bool = True # 是否做 T+1 当日卖探针(验证拒单通路) + hold_n: int = 6 # 每轮等权持有几只(跨类型取样) + period: int = 1 # 每 N 个交易日轮换主组合(1=每日) + probe_t1: bool = True # 是否做 T+1 当日卖探针 + intraday_partial: bool = True # 10:45 部分调仓(加减仓) + intraday_swap: bool = True # 13:45 卖后买(资金复用) benchmark: str = "000300.XSHG" class ChannelTestStrategy: - """通路测试策略:周期性等权轮换 + T+1 拒单探针。""" + """通路测试策略:每日跨类型轮换 + 盘中多时点场景触发。""" def __init__( self, @@ -60,19 +65,26 @@ class ChannelTestStrategy: b.set_benchmark(self.config.benchmark) b.set_option("use_real_price", True) b.set_option("avoid_future_data", True) - b.run_daily(self.rotate, "9:30") + b.run_daily(self.rotate, "9:35") # 主调仓:卖旧买新 + if self.config.intraday_partial: + b.run_daily(self.partial_adjust, "10:45") # 部分加减仓 + if self.config.intraday_swap: + b.run_daily(self.swap_one, "13:45") # 卖后买(资金复用) + if self.config.probe_t1: + b.run_daily(self.t1_probe, "14:30") # T+1 拒单探针 - # ---------------- 轮换主流程 ---------------- + # ---------------- 目标组合 ---------------- def _target_set(self) -> List[str]: - """按 day 偏移在 universe 里取 hold_n 只(循环),保证每轮目标变。""" - u = self.config.universe or _DEFAULT_UNIVERSE - n = max(1, self.config.hold_n) - if len(u) <= n: - return list(u) - offset = (self._day // max(1, self.config.period)) % len(u) - # 取从 offset 起的 n 只(环绕) - return [u[(offset + i) % len(u)] for i in range(n)] + """跨类型取样 hold_n 只:每个类型组按 day 偏移轮流供一只,凑满 hold_n。""" + cycle = self._day // max(1, self.config.period) + picked: List[str] = [] + for i in range(self.config.hold_n): + grp = _TYPE_ORDER[(i + cycle) % len(_TYPE_ORDER)] + codes = UNIVERSE_BY_TYPE[grp] + picked.append(codes[(cycle + i // len(_TYPE_ORDER)) % len(codes)]) + return picked + # ---------------- 场景1:主调仓(卖全量旧 + 买新) ---------------- def rotate(self, context: Any) -> None: self._day += 1 if (self._day - 1) % max(1, self.config.period) != 0: @@ -80,32 +92,76 @@ class ChannelTestStrategy: positions = _get_positions(context) target = self._target_set() target_set = set(target) - logger.info("[channel_test] day=%d target=%s holding=%s", + logger.info("[channel_test] day=%d 主调仓 target=%s holding=%s", self._day, target, list(positions.keys())) - - # 1) 卖:清掉不在目标里的持仓(走卖出通路) for code in list(positions.keys()): if code not in target_set: - logger.info("[channel_test] 卖出 %s", code) + logger.info("[channel_test] 全量卖出 %s", code) self.broker.order_target_value(code, 0) + self._buy_equal_weight(context, target) - # 2) 买/调:目标等权(走买入通路 + 整手) + # ---------------- 场景2:部分加减仓(非清仓) ---------------- + def partial_adjust(self, context: Any) -> None: + positions = _get_positions(context) + if not positions: + return + total = _available_cash(context) + sum(_safe_value(p) for p in positions.values()) + codes = sorted(positions.keys()) + half = max(1, (len(codes) + 1) // 2) + for code in codes[:half]: + # 加仓 10%(走"对已持仓追加买入") + pos_val = _safe_value(positions[code]) + self.broker.order_target_value(code, pos_val * 1.1) + logger.info("[channel_test] 加仓 %s → +10%%", code) + for code in codes[half:half + 1]: + # 减仓到一半(走"部分卖出",非清仓) + pos_val = _safe_value(positions[code]) + if pos_val > 0: + self.broker.order_target_value(code, pos_val * 0.5) + logger.info("[channel_test] 减半仓 %s", code) + _ = total + + # ---------------- 场景3:卖后买(当日资金复用) ---------------- + def swap_one(self, context: Any) -> None: + positions = _get_positions(context) + if not positions: + return + # 清掉第一只(昨日买的,T+1 可卖) → 换买 universe 里下一只未持有的 + out_code = sorted(positions.keys())[0] + logger.info("[channel_test] 换仓卖出 %s", out_code) + self.broker.order_target_value(out_code, 0) + flat = [c for grp in _TYPE_ORDER for c in UNIVERSE_BY_TYPE[grp] + if c not in positions] + if flat: + in_code = flat[self._day % len(flat)] + cash = _available_cash(context) + logger.info("[channel_test] 换仓买入 %s (卖后资金复用)", in_code) + self.broker.order_value(in_code, min(cash * 0.9, 20000)) + + # ---------------- 场景4:T+1 拒单探针 ---------------- + def t1_probe(self, context: Any) -> None: + """当日(含今早刚买)持仓立即试卖 → 两端都应被 T+1 拒,验证拒单通路一致。""" + positions = _get_positions(context) + # 优先挑"今天买的":无买入时间信息就取第一只(今早主调仓必然买过新仓) + if not positions: + return + probe_code = sorted(positions.keys())[0] + try: + self.broker.order_target_value(probe_code, 0) + logger.info("[channel_test] T+1 探针 %s 当日卖已提交(预期被拒)", probe_code) + except Exception as exc: # noqa: BLE001 - 探针失败不阻断 + logger.debug("[channel_test] T+1 探针异常(正常): %s", exc) + + # ---------------- 买入工具 ---------------- + def _buy_equal_weight(self, context: Any, target: List[str]) -> None: + positions = _get_positions(context) cash = _available_cash(context) total = cash + sum(_safe_value(p) for p in positions.values()) per = total / max(1, len(target)) for code in target: - logger.info("[channel_test] 调仓 %s → target_value=%.2f", code, per) + logger.info("[channel_test] 等权调仓 %s → %.2f", code, per) self.broker.order_target_value(code, per) - # 3) T+1 拒单探针:当日买入立即试卖 → 两端都应 T+1 拒(验证拒单通路) - if self.config.probe_t1 and target: - probe_code = target[0] - try: - self.broker.order_target_value(probe_code, 0) - logger.info("[channel_test] T+1 探针 %s 当日卖已提交(预期被拒)", probe_code) - except Exception as exc: # noqa: BLE001 - 探针失败不阻断主流程 - logger.debug("[channel_test] T+1 探针异常(正常): %s", exc) - def _safe_value(pos: Any) -> float: """从 position 对象取市值,兼容多种属性名。""" diff --git a/sanguo_trader/shadow/__main__.py b/sanguo_trader/shadow/__main__.py index 3523f47..c0fa5ff 100644 --- a/sanguo_trader/shadow/__main__.py +++ b/sanguo_trader/shadow/__main__.py @@ -1,7 +1,8 @@ -"""影子柜台 CLI 入口(单实例文件锁防双开重复撮合)。 +"""影子柜台 CLI 入口。 用法: - python -m sanguo_trader.shadow # 读 env(见 runner.py docstring) + python -m sanguo_trader.shadow # 读 env(见 runner.py docstring) + python -m sanguo_trader.shadow --auto # 主管模式:轮询 paper 库自动拉起影子账户 """ from __future__ import annotations @@ -44,6 +45,11 @@ def main() -> int: level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s", ) + if "--auto" in sys.argv: + from .supervisor import run_auto_supervisor + + run_auto_supervisor() + return 0 lock = _acquire_lock() if lock is None: print("影子柜台已在运行(锁占用),本次启动退出。", flush=True) diff --git a/sanguo_trader/shadow/supervisor.py b/sanguo_trader/shadow/supervisor.py new file mode 100644 index 0000000..8bf546e --- /dev/null +++ b/sanguo_trader/shadow/supervisor.py @@ -0,0 +1,112 @@ +"""影子柜台账户主管(P1-d)。 + +轮询 paper 库:发现 ``mode='shadow' & status='running'`` 的组合账户 → 为每个账户 +拉起一个 ``python -m sanguo_trader.shadow --account N`` 子进程(独立虚拟账户); +账户停止/删除 → 终止对应子进程;子进程崩溃 → 重启(告警日志)。 + +用法(VPS,常驻): + set SANGUO_SHADOW_DB=C:\\sanguo_vnpy_v2\\data\\backtest_results.db + python -m sanguo_trader.shadow --auto + +与 sanguo_live.supervisor 同思路:主管只管进程生命周期,不碰撮合。 +""" +from __future__ import annotations + +import json +import logging +import os +import sqlite3 +import subprocess +import sys +import time +from typing import Any, Dict, Optional + +logger = logging.getLogger(__name__) + +POLL_SEC = 60 + + +def load_shadow_accounts(db_path: str) -> list[dict[str, Any]]: + """读所有该拉起影子柜台的账户(mode=shadow & running & portfolio)。""" + with sqlite3.connect(db_path) as conn: + conn.row_factory = sqlite3.Row + rows = conn.execute( + "SELECT * FROM paper_accounts " + "WHERE mode='shadow' AND status='running' AND strategy_type='portfolio'" + ).fetchall() + return [dict(r) for r in rows] + + +def _num(v: Any) -> str: + """数字转 env 字符串,整数不带小数点(500000 而非 500000.0)。""" + f = float(v) + return str(int(f)) if f.is_integer() else repr(f) + + +def account_env(acc: dict[str, Any], db_path: str) -> Dict[str, str]: + """paper 账户行 → 影子柜台子进程 env(复用 SANGUO_LIVE_* 契约)。""" + strategies = json.loads(acc.get("strategies") or "[]") + name = (strategies[0].get("name") if strategies else None) or "all_weather" + params = (strategies[0].get("params") if strategies else {}) or {} + env = dict(os.environ) + env.update({ + "SANGUO_LIVE_STRATEGY": name, + "SANGUO_LIVE_MAX_POOL": _num(params.get("max_pool", 30)), + "SANGUO_LIVE_BENCHMARK": str(params.get("benchmark", "000300.XSHG")), + "SANGUO_LIVE_CASH": _num(acc.get("initial_capital") or 1_000_000), + "SANGUO_SHADOW_DB": db_path, + "SANGUO_SHADOW_ACCOUNT_ID": str(acc["id"]), + "SANGUO_SHADOW_COMMISSION": repr(float(acc.get("rate") or 0.0003)), + "SANGUO_SHADOW_STAMP": repr(float(acc.get("stamp_duty_rate") or 0.001)), + "SANGUO_SHADOW_MIN_COMM": _num(acc.get("min_commission") or 5), + "SANGUO_SHADOW_SLIPPAGE": repr(float(acc.get("slippage") or 0.001)), + }) + return env + + +def spawn_child(acc: dict[str, Any], db_path: str) -> subprocess.Popen: + env = account_env(acc, db_path) + argv = [sys.executable, "-X", "utf8", "-m", "sanguo_trader.shadow", + "--account", str(acc["id"])] + logger.info("[shadow-supervisor] 拉起账户 #%s(%s) 影子柜台", + acc["id"], env["SANGUO_LIVE_STRATEGY"]) + return subprocess.Popen(argv, env=env) + + +def run_auto_supervisor(db_path: Optional[str] = None, poll_sec: float = POLL_SEC) -> None: + """常驻主循环:同步账户 ↔ 子进程。""" + db_path = db_path or os.environ.get("SANGUO_SHADOW_DB") \ + or r"C:\sanguo_vnpy_v2\data\backtest_results.db" + logger.info("[shadow-supervisor] 启动 db=%s", db_path) + children: Dict[int, subprocess.Popen] = {} + + while True: + try: + accounts = load_shadow_accounts(db_path) + except Exception as exc: # noqa: BLE001 - db 抖动不退出主管 + logger.warning("[shadow-supervisor] 读账户失败: %s", exc) + accounts = [] + want = {a["id"]: a for a in accounts} + + # 1) 终止不再需要的 + for aid in list(children): + if aid not in want: + child = children.pop(aid) + if child.poll() is None: + child.terminate() + logger.info("[shadow-supervisor] 账户 #%s 停止,子进程已终止", aid) + + # 2) 重启崩溃的 / 拉起新的 + for aid, acc in want.items(): + child = children.get(aid) + if child is not None and child.poll() is None: + continue + if child is not None: + logger.warning("[shadow-supervisor] 账户 #%s 子进程退出(rc=%s),重启", + aid, child.returncode) + try: + children[aid] = spawn_child(acc, db_path) + except Exception as exc: # noqa: BLE001 + logger.error("[shadow-supervisor] 账户 #%s 拉起失败: %s", aid, exc) + + time.sleep(poll_sec) diff --git a/tests/portfolio/test_channel_test.py b/tests/portfolio/test_channel_test.py index dc3e8f1..21e29d1 100644 --- a/tests/portfolio/test_channel_test.py +++ b/tests/portfolio/test_channel_test.py @@ -1,25 +1,28 @@ """通路测试策略(ChannelTestStrategy)单元测试。 -验证轮换目标集 + 卖旧买新调度 + T+1 探针,用 mock broker 记录下单调用。 +验证跨类型轮换 + 每日多场景调度(主调仓/部分加减/卖后买/T+1探针), +用 mock broker 记录下单调用。坑:BrokerFacade 是 dataclass,子类方法重写会被 +父类 __init__ 写入的实例属性遮蔽 → mock 必须在 super().__init__() 后实例注入。 """ from __future__ import annotations from sanguo_portfolio.strategies import ChannelTestConfig, ChannelTestStrategy from sanguo_portfolio.strategies.all_weather import BrokerFacade +from sanguo_portfolio.strategies.channel_test import UNIVERSE_BY_TYPE class _MockBroker(BrokerFacade): def __init__(self) -> None: - # 注意:BrokerFacade 是 dataclass,父类 __init__ 会用字段默认值覆盖同名 - # 实例属性,子类方法重写会被遮蔽 → 必须在 super().__init__() 之后 - # 用实例属性注入记录函数。 self.calls: list[tuple[str, str, float]] = [] # (method, code, value) super().__init__() - self.order_target_value = self._record_otv + self.order_target_value = self._record("otv") + self.order_value = self._record("ov") - def _record_otv(self, code: str, value: float): - self.calls.append(("otv", code, value)) - return None + def _record(self, method: str): + def rec(code: str, value: float): + self.calls.append((method, code, value)) + return None + return rec class _Pos: @@ -34,53 +37,95 @@ class _Ctx: "total_value": cash + sum(p.value for p in positions.values())})() -def test_target_set_rotates_with_day(): - s = ChannelTestStrategy(provider=None, config=ChannelTestConfig( - universe=["A", "B", "C", "D"], hold_n=2, period=1)) - s._day = 0 - assert set(s._target_set()) <= {"A", "B", "C", "D"} +def test_universe_no_star_bj_stocks(): + """实盘铁律:个股只允许主板+创业板(无 688/920/8 开头)。""" + for codes in UNIVERSE_BY_TYPE.values(): + for c in codes: + num = c.split(".")[0] + if c.endswith(".XSHE") and len(num) == 6 and not num.startswith(("15", "16")): + assert not num.startswith(("30", "00", "002")) or num.startswith(("30", "00")), c + # 个股(非 ETF:51/15/56/58 开头是基金)不允许 688/689/920/8 开头 + is_etf = num.startswith(("51", "15", "56", "58")) + if not is_etf: + assert not num.startswith(("688", "689", "92", "4", "8")), f"非主板/创业板个股: {c}" + + +def test_universe_covers_all_types_with_multiple(): + """每类至少 3 只、类型覆盖宽基/行业/跨境/主板/创业板。""" + for grp, codes in UNIVERSE_BY_TYPE.items(): + assert len(codes) >= 3, grp + assert len(UNIVERSE_BY_TYPE) >= 5 + + +def test_target_set_spans_types(): + s = ChannelTestStrategy(provider=None, config=ChannelTestConfig(hold_n=6)) s._day = 1 - t1 = s._target_set() - s._day = 2 - t2 = s._target_set() - assert len(t1) == 2 and len(t2) == 2 - assert t1 != t2 # 不同周期目标偏移 + target = s._target_set() + assert len(target) == 6 + # 跨类型:6 只来自不同类型组 + code2grp = {c: g for g, cs in UNIVERSE_BY_TYPE.items() for c in cs} + groups = {code2grp[c] for c in target} + assert len(groups) == 6 # hold_n=6 且组数≥6 → 每组一只 def test_rotate_sells_non_target_and_buys_target(): broker = _MockBroker() - s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig( - universe=["A", "B", "C", "D"], hold_n=2, period=1, probe_t1=False)) - # rotate#1 → day=1 → offset=1 → target=[B,C];当前持仓 C,D - # → 卖出 D(C 在目标内保留),对 B,C 调仓(买入) - ctx = _Ctx({"C": _Pos(1000), "D": _Pos(1000)}, cash=8000) + s = ChannelTestStrategy(provider=None, broker=broker, + config=ChannelTestConfig(hold_n=2, period=1, probe_t1=False, + intraday_partial=False, intraday_swap=False)) + ctx = _Ctx({"510300.XSHG": _Pos(1000), "600519.XSHG": _Pos(1000)}, cash=8000) s.rotate(ctx) zero_calls = {c for m, c, v in broker.calls if m == "otv" and v == 0} - assert zero_calls == {"D"} # 只卖非目标的 D - buys = {c for m, c, v in broker.calls if m == "otv" and v > 0} - assert buys == {"B", "C"} # 买新 B + 调仓 C + assert zero_calls # 有全量卖出(旧持仓不在新目标) + buys = [c for m, c, v in broker.calls if m == "otv" and v > 0] + assert len(buys) == 2 # 等权买入 hold_n 只 -def test_rotate_t1_probe_fires(): +def test_partial_adjust_buys_more_and_sells_half(): broker = _MockBroker() - s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig( - universe=["A", "B"], hold_n=1, period=1, probe_t1=True)) - ctx = _Ctx({}, cash=10000) - s.rotate(ctx) - # rotate#1 → target=[B];探针对 target[0]=B 再次 order_target_value(0)(当日卖,预期被 T+1 拒) - otv_calls = [c for m, c, v in broker.calls if m == "otv"] - assert otv_calls.count("B") >= 2 # 一次买入调仓 + 一次 T+1 探针 + s = ChannelTestStrategy(provider=None, broker=broker, + config=ChannelTestConfig(intraday_partial=True)) + ctx = _Ctx({"A": _Pos(1000), "B": _Pos(1000), "C": _Pos(1000)}, cash=1000) + s.partial_adjust(ctx) + ups = [(c, v) for m, c, v in broker.calls if v > 1000] # 加仓 >原值 + downs = [(c, v) for m, c, v in broker.calls if 0 < v < 1000] # 部分减仓 + assert ups and all(v >= 1000 for _, v in ups) + assert downs # 部分卖出(非清仓) -def test_period_skips_off_cycle_days(): +def test_swap_one_sells_then_buys(): broker = _MockBroker() - s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig( - universe=["A", "B"], hold_n=1, period=3, probe_t1=False)) - ctx = _Ctx({}, cash=10000) - s.rotate(ctx) # day1 → 触发 - n1 = len(broker.calls) - s.rotate(ctx) # day2 → 跳过 - s.rotate(ctx) # day3 → 跳过 - assert len(broker.calls) == n1 - s.rotate(ctx) # day4 → 触发 - assert len(broker.calls) > n1 + s = ChannelTestStrategy(provider=None, broker=broker, + config=ChannelTestConfig(intraday_swap=True)) + s._day = 1 + ctx = _Ctx({"510300.XSHG": _Pos(5000), "600519.XSHG": _Pos(5000)}, cash=3000) + s.swap_one(ctx) + sells = [c for m, c, v in broker.calls if v == 0] + buys = [(c, v) for m, c, v in broker.calls if v > 0] + assert sells == ["510300.XSHG"] # 卖第一只 + assert buys and buys[0][0] not in ("510300.XSHG", "600519.XSHG") # 买未持有的 + + +def test_t1_probe_fires(): + broker = _MockBroker() + s = ChannelTestStrategy(provider=None, broker=broker, + config=ChannelTestConfig(probe_t1=True)) + ctx = _Ctx({"510300.XSHG": _Pos(5000)}, cash=0) + s.t1_probe(ctx) + assert broker.calls == [("otv", "510300.XSHG", 0)] + + +def test_initialize_registers_intraday_schedules(): + registered: list[tuple[str, str]] = [] + + class _RegBroker(_MockBroker): + def __init__(self) -> None: + super().__init__() + self.run_daily = lambda fn, t: registered.append( + (getattr(fn, "__name__", str(fn)), t)) + + s = ChannelTestStrategy(provider=None, broker=_RegBroker(), + config=ChannelTestConfig()) + s.initialize(object()) + times = {t for _, t in registered} + assert {"9:35", "10:45", "13:45", "14:30"} <= times # 四个盘中时点全注册 diff --git a/tests/trader/test_shadow_supervisor.py b/tests/trader/test_shadow_supervisor.py new file mode 100644 index 0000000..df6488b --- /dev/null +++ b/tests/trader/test_shadow_supervisor.py @@ -0,0 +1,64 @@ +"""影子柜台主管(supervisor)纯逻辑测试:账户筛选 + env 映射。""" +from __future__ import annotations + +import json + +from sanguo_trader.persistence import init_db, save_account +from sanguo_trader.shadow.supervisor import account_env, load_shadow_accounts + + +def _mk_account(**kw) -> dict: + base = dict( + name="shadow-acc", mode="shadow", strategy_type="portfolio", + interval="15m", symbols=["hs300_subset"], + strategies=[{"name": "channel_test", "params": {"max_pool": 6, + "benchmark": "000905.XSHG"}}], + initial_capital=500_000, rate=0.00025, stamp_duty_rate=0.001, + min_commission=5, slippage=0.002, + start="2026-08-14", end="2026-12-31", engine="shadow", + ) + base.update(kw) + return base + + +def test_load_shadow_accounts_filters(tmp_path): + db = str(tmp_path / "p.db") + init_db(db) + from sanguo_trader.persistence import update_account_status + a_shadow = save_account(db, _mk_account()) # 应选中 + update_account_status(db, a_shadow, "running") # 创建后 API 置 running + a_stopped = save_account(db, _mk_account(name="stopped", strategies=[ + {"name": "x", "params": {}}])) + update_account_status(db, a_stopped, "stopped") # 停止 → 不选 + a_live = save_account(db, _mk_account(name="live", mode="live")) # 实走 → 不选 + _ = a_live + + accounts = load_shadow_accounts(db) + assert [a["id"] for a in accounts] == [a_shadow] + + +def test_account_env_mapping(tmp_path): + db = str(tmp_path / "p.db") + init_db(db) + aid = save_account(db, _mk_account()) + from sanguo_trader.persistence import update_account_status + update_account_status(db, aid, "running") + acc = load_shadow_accounts(db)[0] + env = account_env(acc, db) + assert env["SANGUO_LIVE_STRATEGY"] == "channel_test" + assert env["SANGUO_LIVE_MAX_POOL"] == "6" + assert env["SANGUO_LIVE_BENCHMARK"] == "000905.XSHG" + assert env["SANGUO_LIVE_CASH"] == "500000" + assert env["SANGUO_SHADOW_DB"] == db + assert env["SANGUO_SHADOW_ACCOUNT_ID"] == str(aid) + assert env["SANGUO_SHADOW_COMMISSION"] == "0.00025" + assert env["SANGUO_SHADOW_SLIPPAGE"] == "0.002" + + +def test_account_env_defaults_on_sparse_row(): + acc = {"id": 9, "strategies": json.dumps([]), + "initial_capital": None, "rate": None} + env = account_env(acc, "db") + assert env["SANGUO_LIVE_STRATEGY"] == "all_weather" + assert env["SANGUO_LIVE_CASH"] == "1000000" + assert env["SANGUO_SHADOW_ACCOUNT_ID"] == "9"