This commit is contained in:
2026-09-19 19:45:43 +08:00
parent 7183cb45f8
commit 8131b158b4
60 changed files with 7669 additions and 909 deletions

6
labs/tests/__init__.py Normal file
View File

@@ -0,0 +1,6 @@
"""labs 的测试包。
测试模块内部使用 ``from tests.zt_harness import ...`` 这类绝对导入,
因此 ``labs`` 必须作为顶层包导入、``labs/tests`` 必须是包。
统一入口见 ``labs/run_tests.py``。
"""

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

View File

@@ -1,4 +1,4 @@
"""ETF 离线回归:指标、真实防飞刀/网格算法、限仓、回报持久化。"""
"""ETF 离线回归:指标、分格网格(每格独立买卖)、费用门槛、限仓、回报持久化。"""
from datetime import date, datetime, timedelta
import httpx
@@ -12,12 +12,14 @@ from strategy.etf.config import ETFConfig, load
from strategy.etf.data import DAILY_URL, daily_bars, parse_daily
from strategy.etf.engine import Engine
from strategy.etf.indicators import Indicators, calculate
from strategy.etf.state import Store
from strategy.etf.state import GridLot, Store, SymbolState
CODE, OTHER = '510300.SH', '159915.SZ'
NOW = datetime(2026, 9, 16, 10)
IND = Indicators('20260915', 10, 0.2, 9.5, 10, 10.5, 0.2)
# ma60=10、grid=0.2,指标入场门槛 9.5;引擎再叠加 ma60-1格 = 9.8,最终入场价 9.5
IND = Indicators('20260915', 10, 0.2, 9.5, 10, 10.5, 0.2, entry_band=9.5,
donchian_lo=9.1, donchian_hi=10.9)
def tick(price, now=NOW):
@@ -29,201 +31,414 @@ def position(volume=0, cost=0, available=None):
can_use_volume=volume if available is None else available)
def held(lots: dict, anchor: float, opened_days_ago=5):
"""构造一个已持仓的网格lots = {档位: (数量, 成本)}。"""
state = SymbolState(anchor=anchor)
day = NOW.date() - timedelta(days=opened_days_ago)
for level, (volume, cost) in lots.items():
state.lots[str(level)] = GridLot(volume=volume, cost=cost, bought=day, buys=1)
state.last_buy = max(cost for _, cost in lots.values())
return state
class ETFTests(unittest.TestCase):
"""分格网格:每格独立买入,达到自身目标价就卖掉该格。"""
def setUp(self):
temp = tempfile.TemporaryDirectory()
self.addCleanup(temp.cleanup)
self.path = Path(temp.name) / 'state.json'
self.client = Mock()
self.client.passorder.return_value = {'status': 'success'}
self.cfg = ETFConfig(codes=(CODE,), min_commission=0, commission_rate=0)
# 测试用零佣金、零金额门槛、每格 1 手、3 档:真实默认值为最低佣金 5 元设计。
self.cfg = ETFConfig(codes=(CODE,), buy_hands=1, grid_levels=2, max_hands=3,
min_commission=0, commission_rate=0, min_order_value=0)
self.store = Store(self.path, 'test')
self.engine = Engine(self.client, self.cfg, self.store, 0.1)
def run_price(self, price, pos=None, orders=(), cash=10000, now=NOW, ind=IND):
portfolio = Portfolio(Assets(total=10000, available=cash), {CODE: pos or position()}, list(orders))
self.engine.run(portfolio, {CODE: tick(price, now)}, {CODE: ind}, now)
# ------------------------------------------------------------- 辅助
def run_price(self, price, pos=None, orders=(), cash=10000, now=NOW, ind=IND, engine=None):
portfolio = Portfolio(Assets(total=10000, available=cash),
{CODE: pos or position()}, list(orders))
(engine or self.engine).run(portfolio, {CODE: tick(price, now)}, {CODE: ind}, now)
def buy(self):
self.run_price(9.4)
self.run_price(9.46)
self.client.passorder.assert_called_once()
def keep(self, low, high):
"""让 DipWatch 从 low 反弹到 high反弹幅度必须 ≥ 0.61%)。"""
self.run_price(low)
self.run_price(round(low * 1.007, 3))
self.run_price(high)
def report(self, status=56, filled=100, side=23, price=9.46):
def anchor_grid(self, low=9.3, high=9.38):
"""跌到 low 后反弹到 high≥0.61%)确认,锚点取确认价 high挂单挂在 low。"""
self.run_price(low)
self.run_price(high)
return self.newest()
def ack(self, pending, pos, price=None, side=None, cost=None, status=56, filled=None):
"""把一笔委托做成终态回报,并给出成交后的持仓快照。"""
volume = pending['volume'] if filled is None else filled
side = side or pending['side']
fill_price = pending['price'] if price is None else price
report = OrderItem(stock_code=CODE, remark=pending['id'] + '|etf',
offset_flag={'BUY': 23, 'SELL': 24}[side],
volume_traded=volume, volume_total_original=pending['volume'],
traded_price=fill_price, order_status=status)
after = PositionItem(stock_code=CODE,
volume=pos.volume + (volume if side == 'BUY' else -volume),
open_price=cost if cost is not None else pos.open_price,
can_use_volume=max(0, pos.can_use_volume - volume)
if side == 'SELL' else pos.can_use_volume)
return report, after
def fills(self, volume, cost, available=None):
return {level: (volume, cost, available) for level in self.store.get(CODE).lots}
def newest(self, side='BUY'):
pending = self.store.get(CODE).pending
return OrderItem(stock_code=CODE, remark=pending['id'] + '|etf', offset_flag=side,
volume_traded=filled, volume_total_original=pending['volume'],
traded_price=price, order_status=status)
rows = [p for p in pending.values() if p['side'] == side]
return max(rows, key=lambda p: p['id'])
def test_boll_lower_requires_rebound_and_uses_fixed_limit_order(self):
self.run_price(9.8)
self.run_price(9.4)
self.run_price(9.3)
self.run_price(9.35)
def report(self, pending, status=56, filled=None, price=None, side=None):
volume = pending['volume'] if filled is None else filled
return OrderItem(stock_code=CODE, remark=pending['id'] + '|etf',
offset_flag={'BUY': 23, 'SELL': 24}[side or pending['side']],
volume_traded=volume, volume_total_original=pending['volume'],
traded_price=pending['price'] if price is None else price,
order_status=status)
# ------------------------------------------------------------- 入场
def test_entry_band_requires_rebound_then_ladder_order(self):
"""进入入场区只是开始观察;反弹确认后以 t0 价为锚点挂出锚点档。"""
self.run_price(9.8) # 未进入入场区
self.run_price(9.4) # 进入入场区,开始观察
self.run_price(9.3) # 刷新低点 t0=9.3
self.run_price(9.35) # 反弹不足 0.61%
self.client.passorder.assert_not_called()
self.run_price(9.36)
self.run_price(9.38) # 反弹 (9.38-9.3)/9.3 = 0.86%,锚点取确认价
request = self.client.passorder.call_args.kwargs
self.assertEqual((request['volume'], request['price'], request['pr_type']), (100, 9.36, 11))
self.assertEqual((request['volume'], request['price'], request['pr_type']), (100, 9.38, 11))
self.assertEqual(request['strategy_name'], 'etf')
self.assertTrue(self.store.get(CODE).pending)
self.assertEqual(self.store.get(CODE).anchor, 9.38)
self.assertEqual(set(self.store.get(CODE).pending), {'0'})
def test_leaving_entry_zone_restarts_observation(self):
"""价格弹回入场区上方后,旧低点作废,必须重新形成低点再确认。"""
self.run_price(9.4)
self.run_price(9.3) # t0=9.3
self.run_price(9.8) # 离开入场区,观察点作废
self.run_price(9.38) # 反弹不再基于 9.3
self.client.passorder.assert_not_called()
def test_ladder_places_one_order_per_level_below_anchor(self):
"""锚点 9.4、格距 0.2:价格每跌一格补一档;同一时刻只允许一笔在途委托。"""
cash = 50000
pending = self.anchor_grid(low=9.3, high=9.4) # 锚点 = 9.4
self.assertEqual((pending['level'], pending['price']), (0, 9.4))
report, after = self.ack(pending, position(), price=9.4)
self.run_price(9.4, after, [report], cash=cash) # 锚点档成交
self.assertEqual(self.store.get(CODE).lots['0'].volume, 100)
self.run_price(9.2, position(100, 9.4, 0), cash=cash) # 跌到第 1 档
self.assertEqual(self.store.get(CODE).pending['1']['price'], 9.2) # 锚点-1格
self.run_price(9.2, position(100, 9.4, 0), cash=cash) # 同档不重复挂单
self.assertEqual(sorted(self.store.get(CODE).pending), ['1'])
pending = self.store.get(CODE).pending['1']
report, after = self.ack(pending, position(100, 9.4, 0), price=9.2)
self.run_price(9.0, after, [report], cash=cash) # 第 1 档成交
self.assertEqual(self.store.get(CODE).lots['1'].volume, 100)
self.run_price(9.0, position(200, 9.3, 0), cash=cash) # 跌到第 2 档
self.assertEqual(self.store.get(CODE).pending['2']['price'], 9.0) # 锚点-2格
def test_per_level_fill_cost_and_target(self):
pending = self.anchor_grid()
report, after = self.ack(pending, position(), price=9.28)
self.run_price(9.28, after, [report])
state = self.store.get(CODE)
self.assertEqual(state.lots['0'].volume, 100)
self.assertAlmostEqual(state.lots['0'].cost, 9.28)
self.assertEqual(state.lots['0'].bought, NOW.date())
# 目标价 = 成本 + 格距×2 = 9.28 + 0.4
self.assertAlmostEqual(self.engine.sell_target(9.28, IND), 9.68)
# ------------------------------------------------------------- 卖出
def test_each_level_sells_independently_at_own_target(self):
"""两档成本不同,各自到价才卖,且只卖该档。"""
state = held({0: (100, 9.0), 1: (100, 10.0)}, anchor=10.0)
pos = position(200, 9.5, 200)
self.run_price(9.4, pos, engine=self.engine_with(state))
request = self.client.passorder.call_args.kwargs
self.assertEqual((request['op_type'], request['volume']), (24, 100))
self.assertEqual(self.store.get(CODE).pending['0']['price'], 9.4)
self.assertNotIn('1', self.store.get(CODE).pending)
def engine_with(self, state):
"""把预置状态挂到引擎上,跳过与券商快照的首次核对。"""
self.store.symbols[CODE] = state
engine = Engine(self.client, self.cfg, self.store, 0.1)
return engine
def test_sell_waits_for_target_then_clears_only_that_level(self):
state = held({0: (100, 9.0)}, anchor=9.0)
engine = self.engine_with(state)
pos = position(100, 9.0, 100)
target = engine.sell_target(9.0, IND)
self.run_price(target - 0.001, pos, engine=engine)
self.client.passorder.assert_not_called()
self.run_price(target, pos, engine=engine)
request = self.client.passorder.call_args.kwargs
self.assertEqual((request['op_type'], request['volume']), (24, 100))
pending = self.store.get(CODE).pending['0']
report, after = self.ack(pending, pos)
self.run_price(target, after, [report], engine=engine)
state = self.store.get(CODE)
self.assertEqual(state.lots, {})
self.assertEqual(state.pending, {})
def test_sold_level_is_rebought_when_price_returns(self):
state = held({0: (100, 9.0)}, anchor=9.0)
engine = self.engine_with(state)
pos = position(100, 9.0, 100)
target = engine.sell_target(9.0, IND)
self.run_price(target, pos, engine=engine)
report, after = self.ack(self.store.get(CODE).pending['0'], pos)
self.run_price(target, after, [report], engine=engine)
self.assertEqual(self.store.get(CODE).lots, {})
# 价格回到锚点下方:同一档重新挂买单
self.run_price(9.0, position(), engine=engine)
pending = self.store.get(CODE).pending
self.assertIn('0', pending)
self.assertEqual(pending['0']['side'], 'BUY')
def test_t_plus_one_lot_is_not_sold_same_day(self):
state = held({0: (100, 9.0)}, anchor=9.0, opened_days_ago=0)
self.run_price(10.7, position(100, 9.0, 100), engine=self.engine_with(state))
self.client.passorder.assert_not_called()
def test_partial_available_volume_limits_sell_to_whole_lots(self):
state = held({0: (300, 9.0)}, anchor=9.0)
self.run_price(10.7, position(300, 9.0, 250), engine=self.engine_with(state))
request = self.client.passorder.call_args.kwargs
self.assertEqual((request['op_type'], request['volume']), (24, 200))
def test_max_hold_days_clears_only_stale_level(self):
cfg = ETFConfig(codes=(CODE,), buy_hands=1, grid_levels=2, max_hands=3,
min_commission=0, commission_rate=0, min_order_value=0,
max_hold_days=3)
self.engine = Engine(self.client, cfg, self.store, 0.1)
state = held({0: (100, 9.0)}, anchor=9.0)
self.store.symbols[CODE] = state
with self.assertLogs(level='WARNING'):
self.run_price(9.0, position(100, 9.0, 100))
request = self.client.passorder.call_args.kwargs
self.assertEqual((request['op_type'], request['volume']), (24, 100))
# ------------------------------------------------------------- 核对
def test_fill_waits_for_position_snapshot(self):
pending = self.anchor_grid()
report = self.report(pending)
with self.assertLogs(level='WARNING'):
self.run_price(9.28, orders=[report])
self.assertIn('0', self.store.get(CODE).pending)
report, after = self.ack(pending, position(), price=9.28)
self.run_price(9.28, after, [report])
state = self.store.get(CODE)
self.assertEqual(state.pending, {})
self.assertEqual(state.lots['0'].volume, 100)
self.assertAlmostEqual(state.last_buy, 9.28)
self.client.passorder.assert_called_once()
def test_pending_written_before_network_and_retained_after_timeout(self):
"""提交前先落盘意图;网络异常不解除锁,重启后仍保留待确认状态。"""
def submit(**kwargs):
saved = Store(self.path, 'test').get(CODE).pending
self.assertEqual(saved['id'], kwargs['order_id'])
self.assertEqual({p['id'] for p in saved.values()}, {kwargs['order_id']})
raise TimeoutError('unknown result')
self.client.passorder.side_effect = submit
self.run_price(9.4)
engine = Engine(self.client, self.cfg, Store(self.path, 'test'), 0.1)
engine.orders.busy_cache.clear()
self.run_price(9.3, engine=engine) # 进入入场区
with self.assertLogs(level='ERROR'):
self.run_price(9.46)
self.engine = Engine(self.client, self.cfg, Store(self.path, 'test'), 0.1)
self.run_price(9.38, engine=engine) # 反弹确认,提交时网络异常
self.assertTrue(Store(self.path, 'test').get(CODE).pending)
engine = Engine(self.client, self.cfg, Store(self.path, 'test'), 0.1)
engine.orders.busy_cache.clear()
with self.assertLogs(level='WARNING'):
self.run_price(9.3, now=NOW + timedelta(minutes=10))
self.run_price(9.3, engine=engine, now=NOW + timedelta(minutes=10))
self.client.passorder.assert_called_once()
def test_filled_order_waits_for_position_snapshot(self):
self.buy()
report = self.report()
with self.assertLogs(level='WARNING'):
self.run_price(9.2, orders=[report])
self.assertTrue(self.store.get(CODE).pending)
self.run_price(9.2, position(100, 9.46, 0), [report])
self.assertFalse(self.store.get(CODE).pending)
self.assertEqual(self.store.get(CODE).last_buy, 9.46)
self.client.passorder.assert_called_once()
def test_rejected_order_releases_level_and_does_not_advance(self):
pending = self.anchor_grid()
report = self.report(pending, status=57, filled=0)
self.run_price(9.39, orders=[report]) # 价格已回到锚点上方,不会重挂
state = self.store.get(CODE)
self.assertEqual(state.pending, {})
self.assertEqual(state.lots, {})
self.assertEqual(state.last_buy, 0.0)
# 网格仍保留锚点,价格回到锚点档可重新挂单
self.run_price(9.38)
self.assertIn('0', self.store.get(CODE).pending)
def test_add_requires_another_grid_below_actual_fill(self):
self.buy()
report = self.report()
pos = position(100, 9.46, 0)
self.run_price(9.4, pos, [report])
self.run_price(9.46, pos)
self.client.passorder.assert_called_once()
self.run_price(9.1, pos)
self.run_price(9.16, pos)
self.assertEqual(self.client.passorder.call_count, 2)
def test_external_position_change_rebuilds_grid_from_broker_snapshot(self):
self.run_price(10.0, position(500, 9.9, 500))
state = self.store.get(CODE)
self.assertEqual(state.volume, 500)
self.assertEqual(state.anchor, 9.9)
self.assertEqual(state.lots['0'].volume, 500)
def test_partial_cancel_records_actual_fill_and_never_exceeds_cap(self):
self.cfg = ETFConfig(codes=(CODE,), buy_hands=2, min_commission=0, commission_rate=0)
self.engine = Engine(self.client, self.cfg, self.store, 0.1)
pos = position(800, 10)
self.run_price(9.4, pos)
self.run_price(9.46, pos)
report = self.report(status=53, filled=100)
self.run_price(9.1, position(900, 9.94), [report])
self.run_price(9.16, position(900, 9.94))
self.client.passorder.assert_called_once()
self.assertFalse(self.store.get(CODE).pending)
def test_full_position_blocks_buy_and_zero_position_is_not_a_warning(self):
self.run_price(9.4, position(1000, 10))
self.run_price(9.46, position(1000, 10))
self.client.passorder.assert_not_called()
with self.assertNoLogs(level='WARNING'):
self.run_price(9.8, position())
def test_zero_position_with_retained_broker_cost_can_reopen(self):
self.run_price(9.4, position(0, 10))
self.run_price(9.46, position(0, 10))
self.client.passorder.assert_called_once()
def test_rejected_order_is_logged_and_does_not_advance_anchor(self):
self.buy()
report = self.report(status=57, filled=0)
self.run_price(9.8, orders=[report])
self.assertEqual(self.store.get(CODE).last_buy, 0)
self.assertFalse(self.store.get(CODE).pending)
def test_t_plus_one_tracks_peak_but_only_sells_available_whole_lots(self):
self.run_price(10.7, position(200, 10, 0))
self.run_price(10.55, position(200, 10, 0))
self.client.passorder.assert_not_called()
self.engine = Engine(self.client, self.cfg, Store(self.path, 'test'), 0.1)
self.run_price(10.55, position(200, 10, 100))
order = self.client.passorder.call_args.kwargs
self.assertEqual((order['op_type'], order['volume']), (24, 100))
def test_drop_below_activation_price_still_triggers_profitable_retreat(self):
self.run_price(10.7, position(100, 10))
self.run_price(10.4, position(100, 10))
self.assertEqual(self.client.passorder.call_args.kwargs['op_type'], 24)
def test_cost_change_and_flat_position_reset_peak(self):
self.run_price(10.7, position(100, 10))
self.assertTrue(self.store.get(CODE).armed)
self.run_price(10.5, position(200, 10.4))
self.assertFalse(self.store.get(CODE).armed)
self.run_price(9.8, position())
self.assertEqual(self.store.get(CODE).last_buy, 0)
self.client.passorder.assert_not_called()
def test_fee_floor_prevents_loss_after_commission(self):
cfg = ETFConfig(codes=(CODE,), min_commission=50, commission_rate=0)
self.engine = Engine(self.client, cfg, self.store, 0.1)
self.run_price(10.7, position(100, 10))
self.run_price(10.55, position(100, 10))
def test_full_position_blocks_further_levels(self):
state = self.store.get(CODE)
state.adopt(300, 9.0, NOW.date() - timedelta(days=5))
self.run_price(8.0, position(300, 9.0, 300))
self.client.passorder.assert_not_called()
def test_cash_reserve_and_fixed_lot_no_downsizing(self):
self.run_price(9.4, cash=1900)
self.run_price(9.46, cash=1900)
self.run_price(9.4, cash=1, now=NOW)
self.run_price(9.3, cash=1, now=NOW)
self.client.passorder.assert_not_called()
def test_min_order_value_skips_small_level_orders(self):
cfg = ETFConfig(codes=(CODE,), buy_hands=1, grid_levels=2, max_hands=3,
min_commission=5.0, commission_rate=0.0003, min_order_value=2000)
self.engine = Engine(self.client, cfg, self.store, 0.1)
with self.assertLogs(level='INFO'):
self.run_price(9.4)
self.run_price(9.3)
self.client.passorder.assert_not_called()
def test_multiple_symbols_share_one_cash_budget(self):
cfg = ETFConfig(codes=(CODE, OTHER), min_commission=0, commission_rate=0)
cfg = ETFConfig(codes=(CODE, OTHER), buy_hands=1, grid_levels=2, max_hands=3,
min_commission=0, commission_rate=0, min_order_value=0)
self.engine = Engine(self.client, cfg, self.store, 0.1)
portfolio = Portfolio(Assets(total=10000, available=2500), {}, [])
for price in (9.4, 9.46):
self.engine.run(portfolio, {c: tick(price) for c in cfg.codes}, {c: IND for c in cfg.codes}, NOW)
portfolio = Portfolio(Assets(total=10000, available=2000), {}, [])
for price in (9.3, 9.38):
self.engine.run(portfolio, {c: tick(price) for c in cfg.codes},
{c: IND for c in cfg.codes}, NOW)
self.client.passorder.assert_called_once()
self.assertIn('0', self.store.get(CODE).pending)
self.assertFalse(self.store.get(OTHER).pending)
def test_pending_later_symbol_reserves_cash_before_first_symbol(self):
cfg = ETFConfig(codes=(CODE, OTHER), buy_hands=1, grid_levels=2, max_hands=3,
min_commission=0, commission_rate=0, min_order_value=0)
self.store.get(OTHER).pending['0'] = dict(id='ETF-BUY-pending', side='BUY', volume=100,
base_volume=0, reserved=950, level=0,
price=9.5)
self.engine = Engine(self.client, cfg, self.store, 0.1)
portfolio = Portfolio(Assets(total=10000, available=2000), {}, [])
with self.assertLogs(level='WARNING'):
for price in (9.4, 9.3):
self.engine.run(portfolio, {CODE: tick(price)}, {CODE: IND}, NOW)
self.client.passorder.assert_not_called()
def test_other_strategy_order_blocks_same_symbol_without_cancel(self):
report = OrderItem(stock_code=CODE, remark='TREN-BUY-other', offset_flag=23,
order_status=50, insert_date='20260916', insert_time='093000')
self.run_price(9.4, orders=[report])
self.run_price(9.46, orders=[report])
self.run_price(9.3, orders=[report])
self.client.passorder.assert_not_called()
self.client.cancel_by_id.assert_not_called()
def test_pending_later_symbol_reserves_cash_before_first_symbol(self):
cfg = ETFConfig(codes=(CODE, OTHER), min_commission=0, commission_rate=0)
self.store.get(OTHER).pending = dict(id='ETF-BUY-pending', side='BUY', volume=100,
base_volume=0, reserved=950)
self.engine = Engine(self.client, cfg, self.store, 0.1)
portfolio = Portfolio(Assets(total=10000, available=2500), {}, [])
with self.assertLogs(level='WARNING'):
for price in (9.4, 9.46):
self.engine.run(portfolio, {CODE: tick(price)}, {CODE: IND}, NOW)
self.client.passorder.assert_not_called()
def test_on_road_or_unknown_order_never_opens_another_buy(self):
def test_on_road_or_unknown_order_never_opens_grid(self):
pos = position()
pos.on_road_volume = 100
self.run_price(9.4, pos)
self.run_price(9.46, pos)
self.run_price(9.3, pos)
unknown = OrderItem(stock_code=CODE, order_status=255)
self.run_price(9.4, orders=[unknown])
self.run_price(9.46, orders=[unknown])
self.run_price(9.3, orders=[unknown])
self.client.passorder.assert_not_called()
def test_excluded_symbol_is_neither_bought_nor_sold(self):
self.engine.excluded.add(CODE)
self.run_price(9.4)
self.run_price(9.46)
self.run_price(10.7, position(100, 10))
self.run_price(10.55, position(100, 10))
self.run_price(9.3)
self.run_price(10.7, position(100, 9.0, 100))
self.client.passorder.assert_not_called()
def test_entry_gate_uses_configured_band_not_boll_lower(self):
"""入场门槛取自 Indicators.entry_bandband_type 决定它由谁计算。"""
donchian_ind = Indicators('20260915', 10, 0.2, 9.5, 10, 10.5, 0.2,
entry_band=9.0, donchian_lo=8.8, donchian_hi=11.0)
self.run_price(9.4, ind=donchian_ind)
self.run_price(9.35, ind=donchian_ind) # 高于 Donchian 门槛 9.0
self.client.passorder.assert_not_called()
self.run_price(8.95, ind=donchian_ind) # 进入入场区
self.run_price(8.90, ind=donchian_ind) # 刷新低点
self.run_price(8.96, ind=donchian_ind) # 反弹 0.67% 确认
self.client.passorder.assert_called_once()
self.assertEqual(self.client.passorder.call_args.kwargs['price'], 8.96)
def test_band_type_switch_changes_entry_band(self):
rows = IndicatorTests().bars()
don = calculate(rows, NOW.date(), ETFConfig(codes=(CODE,), band_type='donchian',
donchian_period=20, donchian_pct=15.0))
boll = calculate(rows, NOW.date(), ETFConfig(codes=(CODE,), band_type='boll'))
self.assertEqual((don.donchian_lo, don.donchian_hi), (9.0, 11.0))
# 常数序列下 ma60=10、grid=2两条通道都被 ma60-1格 压到 8.0
# 因此这里验证通道本身确实换了,而不是只看最终门槛。
self.assertEqual(don.lower, boll.lower)
self.assertAlmostEqual(boll.entry_band, min(boll.lower, boll.ma60 - boll.grid))
widened = calculate(rows, NOW.date(), ETFConfig(codes=(CODE,), band_type='donchian',
donchian_period=20, donchian_pct=5.0))
self.assertAlmostEqual(widened.entry_band, min(9.0 + (11.0 - 9.0) * 0.05,
widened.ma60 - widened.grid))
def test_invalid_or_stale_tick_cannot_trade(self):
for t in (None, Tick(10), tick(float('nan')), tick(9.4, NOW - timedelta(days=1)),
tick(9.4, NOW - timedelta(seconds=91))):
self.assertFalse(self.engine.fresh_tick(t, NOW))
self.assertTrue(self.engine.fresh_tick(tick(9.4), NOW))
# ------------------------------------------------------------- 状态
def test_state_roundtrip_keeps_lots_and_pending(self):
state = self.store.get(CODE)
state.anchor = 9.3
state.last_buy = 9.28
state.lot(0).volume = 200
state.lot(0).cost = 9.28
state.lot(0).bought = NOW.date()
state.lot(0).buys = 1
state.pending['1'] = dict(id='ETF-BUY-abc', side='BUY', volume=100, base_volume=200,
reserved=913.0, level=1, price=9.1)
self.store.save()
again = Store(self.path, 'test').get(CODE)
self.assertEqual(again.anchor, 9.3)
self.assertEqual(again.volume, 200)
self.assertEqual(again.lots['0'].bought, NOW.date())
self.assertEqual(again.pending['1']['level'], 1)
def test_v1_state_is_migrated_to_anchor_lot(self):
self.path.write_text(
'{"version": 1, "account": "test", "symbols": {"510300.SH": '
'{"volume": 200, "cost": 9.9, "last_buy": 9.8, "armed": true, "sell_grid": 0.2, '
'"peak": 3, "hold_days": 4, "pending": {}}}}', encoding='utf-8')
state = Store(self.path, 'test').get(CODE)
self.assertEqual(state.volume, 200)
self.assertEqual(state.anchor, 9.9)
self.assertEqual(state.lots['0'].volume, 200)
self.assertEqual(state.pending, {})
def test_corrupt_state_does_not_silently_start_empty(self):
self.path.write_text('{', encoding='utf-8')
for payload in ('{', '{"version": 9, "account": "test", "symbols": {}}',
'{"version": 2, "account": "other", "symbols": {}}',
'{"version": 2, "account": "test", "symbols": {"510300.SH": '
'{"anchor": 9.0, "lots": {"0": {"volume": 100, "cost": 0}}, "pending": {}}}}',
'{"version": 2, "account": "test", "symbols": {"510300.SH": '
'{"anchor": 0, "lots": {"0": {"volume": 100, "cost": 9.0}}, "pending": {}}}}',
'{"version": 2, "account": "test", "symbols": {"510300.SH": '
'{"anchor": 9.0, "lots": {"0": {"volume": 50, "cost": 9.0}}, "pending": {}}}}'):
with self.subTest(payload=payload):
self.path.write_text(payload, encoding='utf-8')
with self.assertRaises(ValueError):
Store(self.path, 'test')
def test_engine_rejects_symbol_removed_with_pending_order(self):
state = self.store.get(CODE)
state.pending['0'] = dict(id='ETF-BUY-x', side='BUY', volume=100, base_volume=0,
reserved=950.0, level=0, price=9.5)
with self.assertRaises(ValueError):
Store(self.path, 'test')
Engine(self.client, ETFConfig(codes=(OTHER,)), self.store, 0.1)
class IndicatorTests(unittest.TestCase):
@@ -241,6 +456,7 @@ class IndicatorTests(unittest.TestCase):
rows = self.bars() + [dict(date='20260916', high=999, low=1, close=999)]
ind = calculate(rows, NOW.date(), cfg)
self.assertEqual((ind.ma60, ind.atr, ind.lower, ind.upper, ind.grid), (10, 2, 10, 10, 2))
self.assertAlmostEqual(ind.entry_band, min(ind.lower, ind.ma60 - ind.grid))
def test_atr_accounts_for_gap_and_uses_wilder_smoothing(self):
rows = self.bars()
@@ -267,32 +483,41 @@ class ConfigAndDataTests(unittest.TestCase):
def test_config_rejects_excess_hands_and_invalid_codes(self):
for kwargs in ({'max_hands': 11}, {'buy_hands': 11}, {'buy_hands': True},
{'atr_multiplier': float('nan')}, {'codes': ('920202.BJ',)},
{'codes': (CODE, CODE)}, {'codes': ()}):
{'codes': (CODE, CODE)}, {'codes': ()}, {'min_hold_days': -1},
{'max_hold_days': -1}, {'grid_levels': 0}, {'grid_levels': 10},
{'band_type': 'ma'}, {'donchian_pct': 0}):
with self.assertRaises(ValueError):
ETFConfig(**dict({'codes': (CODE,)}, **kwargs))
def test_default_file_loads(self):
cfg = load()
self.assertEqual((cfg.buy_hands, cfg.max_hands), (1, 10))
self.assertEqual(cfg.band_type, 'boll')
self.assertGreaterEqual(cfg.grid_levels, 1)
self.assertGreater(cfg.sell_grid_mult, 0)
self.assertGreaterEqual(cfg.min_order_value, 1000)
class DailyDataTests(unittest.TestCase):
def row(self, day=20260915, **changes):
return dict(dict(ts_code=CODE, trade_date=day, open=10, high=11, low=9, close=10), **changes)
def test_request_and_sample_shape(self):
def test_bare_list_response_is_supported(self):
"""线上 /etf/daily 直接返回一维数组(倒序),旧版是 {code, details} 包装。"""
def respond(request):
self.assertEqual(str(request.url), DAILY_URL + '?code=' + CODE)
self.assertNotIn('x-token', request.headers)
return httpx.Response(200, json={'code': 0, 'message': '', 'details': [self.row()]})
return httpx.Response(200, json=[self.row(20260915)])
with httpx.Client(transport=httpx.MockTransport(respond)) as client:
self.assertEqual(daily_bars(client, CODE, NOW.date()),
[dict(date='20260915', open=10.0, high=11.0, low=9.0, close=10.0)])
def test_legacy_envelope_response_still_supported(self):
payload = {'code': 0, 'message': '', 'details': [self.row(20260915)]}
self.assertEqual([b['date'] for b in parse_daily(payload, CODE, NOW.date())], ['20260915'])
def test_sort_filter_then_limit_and_numeric_strings(self):
payload = dict(code=0, details=[self.row(20260916), self.row(20260915, close='10.5'),
self.row(20260914), self.row(20260917)])
payload = [self.row(20260916), self.row(20260915, close='10.5'),
self.row(20260914), self.row(20260917)]
bars = parse_daily(payload, CODE, NOW.date(), count=1)
self.assertEqual([b['date'] for b in bars], ['20260915'])
self.assertEqual(bars[0]['close'], 10.5)
@@ -307,7 +532,8 @@ class DailyDataTests(unittest.TestCase):
def test_bad_business_response_is_rejected(self):
for payload in (None, [], {}, {'code': False, 'details': [self.row()]},
{'code': 1, 'message': 'failed'}, {'code': 0, 'details': []},
{'code': 0, 'details': {}}, {'code': 0, 'details': None}):
{'code': 0, 'details': {}}, {'code': 0, 'details': None},
[self.row(20260916)]):
with self.subTest(payload=payload), self.assertRaises(ValueError):
parse_daily(payload, CODE, NOW.date())
@@ -316,11 +542,11 @@ class DailyDataTests(unittest.TestCase):
[self.row(20260230)], [self.row(close=float('nan'))],
[self.row(open=True)], [self.row(low=12)], [self.row(close=None)]):
with self.subTest(rows=rows), self.assertRaises(ValueError):
parse_daily(dict(code=0, details=rows), CODE, NOW.date())
parse_daily(rows, CODE, NOW.date())
def test_external_history_flows_into_real_indicators(self):
details = [self.row(int(row['date'])) for row in IndicatorTests().bars()]
rows = parse_daily(dict(code=0, details=details), CODE, NOW.date())
rows = parse_daily([self.row(int(row['date'])) for row in IndicatorTests().bars()],
CODE, NOW.date())
ind = calculate(rows, NOW.date(), ETFConfig(codes=(CODE,)))
self.assertEqual((ind.ma60, ind.atr, ind.grid), (10, 2, 2))

View File

@@ -0,0 +1,147 @@
"""ETF 配置:固定文件加载、缺文件返回 None、全局默认与逐标的覆盖的校验。"""
import tempfile
import textwrap
import unittest
from pathlib import Path
from unittest.mock import patch
import yaml
import config
from config import EtfConfig, EtfDefaults, EtfSymbolConfig
GLOBAL = {"qmt_base_url": "unused", "api_host": "unused", "hosts": {"test": "account"}}
SYMBOL = 'symbols: {"510300.SH": {is_t0: false, buy_shares: 1000, atr_multiplier: 1.0, inner_step: 0.7}}\n'
class EtfConfigTests(unittest.TestCase):
def load(self, etf: str | None = None, account: dict | None = None):
"""在临时目录里生成 _global.yaml / account.yaml / _etf.yaml 并加载。"""
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "_global.yaml").write_text(
yaml.safe_dump(dict(GLOBAL, qmt_data_dir=directory)), encoding="utf-8"
)
(root / "account.yaml").write_text(
yaml.safe_dump(account or {"buy_value": 1000, "strategy": "etf"}),
encoding="utf-8",
)
if etf is not None:
(root / "_etf.yaml").write_text(textwrap.dedent(etf), encoding="utf-8")
with patch.object(config, "global_config"), \
patch.object(config, "account_config"), \
patch.object(config, "etf_config"):
config.load(root, "test")
return config.etf_config
def reject(self, etf: str, message: str):
with self.assertRaisesRegex(ValueError, message):
self.load(etf)
# ------------------------------------------------------------- 加载
def test_missing_file_returns_none(self):
"""_etf.yaml 缺失不是错误:其它策略的账户不应因此启动失败。"""
self.assertIsNone(self.load())
def test_loads_defaults_and_symbols(self):
loaded = self.load(
"""
defaults:
atr_period: 10
add_pct: 2.5
symbols:
"159915.SZ":
is_t0: true
buy_shares: 2000
max_shares: 20000
atr_multiplier: 1.5
inner_step: 1.2
"""
)
self.assertIsInstance(loaded, EtfConfig)
self.assertEqual(loaded.codes, ("159915.SZ",))
self.assertEqual(loaded.defaults.atr_period, 10)
self.assertEqual(loaded.defaults.add_pct, 2.5)
symbol = loaded.symbol("159915.SZ")
self.assertTrue(symbol.is_t0)
self.assertEqual(symbol.buy_shares, 2000)
self.assertEqual(symbol.max_shares, 20000)
def test_unspecified_defaults_fall_back_to_dataclass(self):
loaded = self.load(SYMBOL)
self.assertEqual(loaded.defaults, EtfDefaults())
self.assertEqual(loaded.defaults.commission_rate, 0.0003)
def test_symbol_inherits_uncovered_defaults(self):
"""max_shares 缺省按 max_adds + 1 档计算,其余未覆盖项回落全局默认。"""
loaded = self.load(SYMBOL)
symbol = loaded.symbol("510300.SH")
self.assertEqual(symbol.max_shares, 1000 * (loaded.defaults.max_adds + 1))
self.assertEqual(symbol.get("rebound_pct"), loaded.defaults.rebound_pct)
self.assertEqual(symbol.get("inner_grids"), loaded.defaults.inner_grids)
self.assertEqual(symbol.get("max_hold_days"), loaded.defaults.max_hold_days)
self.assertEqual(symbol.get("max_grid_span_pct"), loaded.defaults.max_grid_span_pct)
def test_symbol_override_wins_over_default(self):
loaded = self.load(
'defaults: {rebound_pct: 0.5}\n' + SYMBOL.replace(
"inner_step: 0.7", "inner_step: 0.7, rebound_pct: 0.4"
)
)
self.assertEqual(loaded.symbol("510300.SH").get("rebound_pct"), 0.4)
def test_unknown_symbol_is_rejected_by_lookup(self):
loaded = self.load(SYMBOL)
with self.assertRaisesRegex(KeyError, "159915.SZ"):
loaded.symbol("159915.SZ")
# ------------------------------------------------------------- 校验
def test_unknown_keys_are_rejected(self):
self.reject("foo: 1\n", "foo")
self.reject("defaults: {atr_period: 14, foo: 1}\n" + SYMBOL, "foo")
self.reject(SYMBOL.replace("inner_step: 0.7", "inner_step: 0.7, add_pct: 2"), "add_pct")
def test_symbols_must_be_a_non_empty_mapping(self):
self.reject("defaults: {atr_period: 14}\n", "symbols")
def test_symbol_code_must_be_a_listed_etf(self):
self.reject('symbols: {"600000.SH": {is_t0: false, buy_shares: 100, '
'atr_multiplier: 1.0, inner_step: 0.7}}\n', "600000.SH")
def test_required_symbol_fields_cannot_fall_back(self):
"""is_t0、buy_shares、atr_multiplier、inner_step 没有安全的全局默认值。"""
self.reject('symbols: {"510300.SH": {is_t0: false, buy_shares: 1000, '
'atr_multiplier: 1.0}}\n', "inner_step")
def test_symbol_field_types_are_checked(self):
self.reject(SYMBOL.replace("is_t0: false", "is_t0: 1"), "is_t0")
self.reject(SYMBOL.replace("buy_shares: 1000", "buy_shares: 150"), "100")
self.reject(SYMBOL.replace("buy_shares: 1000", "buy_shares: 0"), "buy_shares")
self.reject(SYMBOL.replace("atr_multiplier: 1.0", "atr_multiplier: true"), "atr_multiplier")
self.reject(SYMBOL.replace("atr_multiplier: 1.0", "atr_multiplier: 0"), "atr_multiplier")
def test_defaults_ranges_are_checked(self):
self.reject("defaults: {max_adds: 10}\n" + SYMBOL, "max_adds")
self.reject("defaults: {max_adds: -1}\n" + SYMBOL, "max_adds")
self.reject("defaults: {max_adds: 1.5}\n" + SYMBOL, "max_adds")
self.reject("defaults: {max_grid_span_pct: 101}\n" + SYMBOL, "max_grid_span_pct")
self.reject("defaults: {atr_period: 0}\n" + SYMBOL, "atr_period")
self.reject("defaults: {min_profit_pct: 0}\n" + SYMBOL, "min_profit_pct")
# 0 表示不止损,是合法值;佣金允许为 0回测口径
loaded = self.load("defaults: {max_hold_days: 0, commission_rate: 0, min_commission: 0}\n" + SYMBOL)
self.assertEqual(loaded.defaults.max_hold_days, 0)
self.assertEqual(loaded.defaults.commission_rate, 0)
def test_rebound_must_be_below_add_pct(self):
"""反弹确认价必须早于补仓触发,否则条件自相矛盾。"""
self.reject("defaults: {rebound_pct: 3.0, add_pct: 3.0}\n" + SYMBOL, "rebound_pct")
def test_symbol_defaults_are_shared_not_copied(self):
loaded = self.load(SYMBOL)
self.assertIs(loaded.symbol("510300.SH").defaults, loaded.defaults)
self.assertFalse(EtfSymbolConfig().is_t0)
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,324 @@
"""ETF 信号层:白名单展开、已收盘日线校验、指标计算与取数缓存。"""
from datetime import date, datetime, timedelta
from decimal import Decimal, ROUND_CEILING
from statistics import fmean
import unittest
from unittest.mock import Mock, patch
import httpx
from config import AccountConfig, EtfConfig, EtfDefaults, EtfSymbolConfig, GlobalConfig
from libs.signal import SignalItem
from strategy.etf import signal as etf_signal
CODE, OTHER = "510300.SH", "159915.SZ"
TODAY = date(2026, 9, 16)
# 最后一根日线为 2026-09-15前一交易日既不过期也不混入当日未收盘日线。
LAST_DAY = date(2026, 9, 15)
def raw_bar(day: str, close: float = 10.0, spread: float = 1.0, code: str = CODE) -> dict:
"""构造接口口径的一行日线:以 close 为中轴,上下各半个 spread。"""
close = float(close)
return dict(
ts_code=code,
trade_date=day,
open=close,
high=close + spread / 2,
low=close - spread / 2,
close=close,
)
def parse(*payload: dict, code: str = CODE, count: int = 120) -> list[dict]:
"""把接口口径的日线交给 parse_daily得到指标层使用的一行含 date"""
return etf_signal.parse_daily(list(payload), code, TODAY, count)
def raw_payload(
count: int = 90, close: float = 10.0, spread: float = 1.0, code: str = CODE
) -> list[dict]:
"""生成截至 LAST_DAY 的 count 根横盘日线(接口口径,升序)。"""
return [
raw_bar(
(LAST_DAY - timedelta(days=count - 1 - index)).strftime("%Y%m%d"),
close,
spread,
code,
)
for index in range(count)
]
def bars(count: int = 90, close: float = 10.0, spread: float = 1.0) -> list[dict]:
"""生成指标层口径(已解析)的 count 根横盘日线。"""
return parse(*raw_payload(count, close, spread))
def symbol(**overrides) -> EtfSymbolConfig:
base = dict(is_t0=False, buy_shares=1000, atr_multiplier=1.0, inner_step=0.7)
base.update(overrides)
return EtfSymbolConfig(**{k: v for k, v in base.items() if v is not None})
def runtime(etf: EtfConfig | None = None, host: str = "http://api.test") -> Mock:
rt = Mock()
rt.etf_cfg = etf if etf is not None else EtfConfig(
defaults=EtfDefaults(), symbols={CODE: symbol()}
)
rt.global_cfg = GlobalConfig(api_host=host)
rt.account_cfg = AccountConfig(strategy="etf")
return rt
class ParseDailyTests(unittest.TestCase):
def test_bare_list_is_supported_and_sorted_ascending(self):
payload = [raw_bar("20260915"), raw_bar("20260912")]
rows = etf_signal.parse_daily(payload, CODE, TODAY)
self.assertEqual([row["date"] for row in rows], ["20260912", "20260915"])
self.assertEqual(rows[0]["close"], 10.0)
def test_legacy_envelope_is_supported(self):
payload = {"code": 0, "message": "", "details": [raw_bar("20260915")]}
self.assertEqual(
[row["date"] for row in etf_signal.parse_daily(payload, CODE, TODAY)],
["20260915"],
)
def test_today_and_future_bars_are_dropped_then_limited(self):
payload = [
raw_bar("20260916"), raw_bar("20260917"),
raw_bar("20260915"), raw_bar("20260914"),
]
rows = etf_signal.parse_daily(payload, CODE, TODAY, count=1)
self.assertEqual([row["date"] for row in rows], ["20260915"])
def test_numeric_strings_are_converted(self):
rows = etf_signal.parse_daily([raw_bar("20260915", close="10.5")], CODE, TODAY)
self.assertEqual(rows[0]["close"], 10.5)
def test_bad_payloads_are_rejected(self):
cases = [
None, [], {}, {"code": 1, "message": "failed"},
{"code": False, "details": [raw_bar("20260915")]},
{"code": 0, "details": []}, {"code": 0, "details": None},
]
for payload in cases:
with self.subTest(payload=payload), self.assertRaises(ValueError):
etf_signal.parse_daily(payload, CODE, TODAY)
def test_bad_rows_are_rejected(self):
cases = [
[raw_bar("20260915", code=OTHER)], # 证券归属不一致
[raw_bar("20260915"), raw_bar("20260915")], # 日期重复
[raw_bar("20260230")], # 非法日期
[raw_bar("20260915", close=float("nan"))], # 非有限
[dict(raw_bar("20260915"), open=True)], # 布尔价格
[dict(raw_bar("20260915"), low=12)], # OHLC 关系异常
[dict(raw_bar("20260915"), close=None)], # 缺失价格
]
for rows in cases:
with self.subTest(rows=rows), self.assertRaises(ValueError):
etf_signal.parse_daily(rows, CODE, TODAY)
def test_count_must_be_positive_integer(self):
for count in (0, -1, 1.5, True):
with self.subTest(count=count), self.assertRaises(ValueError):
etf_signal.parse_daily([raw_bar("20260915")], CODE, TODAY, count)
class CalculateTests(unittest.TestCase):
"""横盘日线ATR = spread、MA60 = close、格距 = max(ATR×倍数, MA60×0.5%, 0.001)。"""
def test_flat_bars_produce_expected_indicators(self):
values = etf_signal.calculate(bars(), symbol(), EtfDefaults(), TODAY)
self.assertEqual(values[etf_signal.IND_MA60], 10.0)
self.assertEqual(values[etf_signal.IND_ATR], 1.0)
self.assertEqual(values[etf_signal.IND_CHANNEL_LOW], 9.5)
self.assertEqual(values[etf_signal.IND_CHANNEL_HIGH], 10.5)
# 入场门槛 = min(9.5 + 1×15%, 10) = 9.65
self.assertAlmostEqual(values[etf_signal.IND_ENTRY], 9.65)
self.assertEqual(values[etf_signal.IND_GRID], 1.0)
self.assertEqual(values[etf_signal.IND_PRICE], 10.0)
self.assertAlmostEqual(values[etf_signal.IND_GRID_PCT], 10.0)
self.assertAlmostEqual(values[etf_signal.IND_ADD_PRICE], 9.7)
def test_atr_multiplier_and_grid_floor(self):
low_atr = etf_signal.calculate(bars(spread=0.02), symbol(atr_multiplier=1.0),
EtfDefaults(), TODAY)
# ATR=0.02 低于 MA60×0.5% = 0.05 的百分比下限
self.assertAlmostEqual(low_atr[etf_signal.IND_GRID], 0.05)
wide = etf_signal.calculate(bars(), symbol(atr_multiplier=0.5), EtfDefaults(), TODAY)
self.assertEqual(wide[etf_signal.IND_GRID], 0.5)
def test_grid_is_rounded_up_to_tick(self):
rows = bars()
for row in rows:
row["close"] += 0.0004
row["open"], row["high"], row["low"] = row["close"], row["high"] + 0.0004, row["low"] + 0.0004
values = etf_signal.calculate(rows, symbol(), EtfDefaults(), TODAY)
raw = values[etf_signal.IND_ATR] * 1.0
expected = float(Decimal(str(raw)).quantize(Decimal("0.001"), rounding=ROUND_CEILING))
self.assertEqual(values[etf_signal.IND_GRID], expected)
self.assertEqual(round(values[etf_signal.IND_GRID], 3), values[etf_signal.IND_GRID])
def test_rising_close_uses_wilder_smoothing(self):
rows = [dict(row) for row in bars()]
for index, row in enumerate(rows):
shift = index * 0.01
row["close"] += shift
row["open"], row["high"], row["low"] = (
row["close"], row["close"] + 0.5, row["close"] - 0.5
)
values = etf_signal.calculate(rows, symbol(), EtfDefaults(), TODAY)
closes = [row["close"] for row in rows]
self.assertAlmostEqual(values[etf_signal.IND_MA60], fmean(closes[-60:]))
self.assertGreaterEqual(values[etf_signal.IND_ATR], 1.0)
def test_insufficient_or_stale_bars_are_rejected(self):
with self.assertRaisesRegex(ValueError, "61"):
etf_signal.calculate(bars(40), symbol(), EtfDefaults(), TODAY)
# 最近日线距今超过 15 个自然日即放弃该标的
with self.assertRaisesRegex(ValueError, "自然日"):
etf_signal.calculate(bars(), symbol(), EtfDefaults(), TODAY + timedelta(days=20))
def test_bad_atr_period_is_rejected(self):
defaults = EtfDefaults()
defaults.atr_period = 1
with self.assertRaisesRegex(ValueError, "atr_period"):
etf_signal.calculate(bars(), symbol(), defaults, TODAY)
class DailyBarsTests(unittest.TestCase):
def test_request_url_and_no_token_header(self):
def respond(request):
self.assertEqual(
str(request.url), etf_signal.DAILY_URL + "?code=" + CODE
)
self.assertNotIn("x-token", request.headers)
return httpx.Response(200, json=[raw_bar("20260915")])
with httpx.Client(transport=httpx.MockTransport(respond)) as client:
rows = etf_signal.daily_bars(client, CODE, TODAY)
self.assertEqual([row["date"] for row in rows], ["20260915"])
def test_custom_endpoint_is_used(self):
def respond(request):
self.assertEqual(str(request.url), "http://api.test/etf/daily?code=" + CODE)
return httpx.Response(200, json=[raw_bar("20260915")])
with httpx.Client(transport=httpx.MockTransport(respond)) as client:
etf_signal.daily_bars(
client, CODE, TODAY, endpoint="http://api.test/etf/daily"
)
def test_http_and_json_errors_propagate(self):
for status, content in ((404, "{}"), (200, "<html>error</html>")):
with httpx.Client(
transport=httpx.MockTransport(lambda r: httpx.Response(status, text=content))
) as client, self.assertRaises((httpx.HTTPStatusError, ValueError)):
etf_signal.daily_bars(client, CODE, TODAY)
def daily_response(payload: object, code: str = CODE) -> httpx.Response:
"""构造带 request 的 200 响应httpx 的 raise_for_status 需要 request。"""
request = httpx.Request("GET", etf_signal.DAILY_URL, params={"code": code})
return httpx.Response(200, json=payload, request=request)
def boom(url=None, params=None):
"""模拟连接失败httpx 只接受真实的 RequestError 子类实例。"""
raise httpx.ConnectError("boom")
class GenSignalsTests(unittest.TestCase):
def setUp(self):
self.client = Mock(spec=httpx.Client)
self.client.get.side_effect = lambda url, params=None: daily_response(
raw_payload(code=params["code"]), params["code"]
)
for cache in (etf_signal._daily_cache, etf_signal._fetched, etf_signal._retry_at):
cache.clear()
patch.object(etf_signal, "_history_client", self.client).start()
self.addCleanup(patch.stopall)
def config(self, codes=(CODE,)):
return EtfConfig(
defaults=EtfDefaults(),
symbols={code: symbol() for code in codes},
)
def test_signals_follow_config_order_and_carry_indicators(self):
signals = etf_signal.gen_signals(runtime(self.config((OTHER, CODE))))
self.assertEqual([item.code for item in signals], [OTHER, CODE])
item = signals[0]
self.assertIsInstance(item, SignalItem)
self.assertEqual(item.signal_key, "etf")
self.assertEqual(item.last_close, 10.0)
self.assertIn("ETF网格", item.desc)
self.assertEqual(item.tech_indicator[etf_signal.IND_GRID], 1.0)
self.assertEqual(item.tech_indicator[etf_signal.IND_MA60], 10.0)
self.assertEqual(item.tech_indicator[etf_signal.IND_ENTRY], 9.65)
def test_missing_etf_config_yields_empty_list(self):
rt = runtime()
rt.etf_cfg = None
self.assertEqual(etf_signal.gen_signals(rt), [])
def test_broken_symbol_is_skipped_without_breaking_others(self):
"""CODE 的日线证券代码不一致时放弃该标的OTHER 正常生成。"""
payloads = {CODE: [raw_bar("20260915", code=OTHER)], OTHER: raw_payload(code=OTHER)}
self.client.get.side_effect = lambda url, params=None: daily_response(
payloads[params["code"]], params["code"]
)
signals = etf_signal.gen_signals(runtime(self.config((CODE, OTHER))))
self.assertEqual([item.code for item in signals], [OTHER])
def test_daily_bars_are_cached_per_code_per_day(self):
rt = runtime(self.config())
etf_signal.gen_signals(rt)
etf_signal.gen_signals(rt)
self.assertEqual(self.client.get.call_count, 1)
def test_failure_is_throttled_by_retry_window(self):
self.client.get.side_effect = boom
rt = runtime(self.config())
self.assertEqual(etf_signal.gen_signals(rt), [])
self.assertEqual(self.client.get.call_count, 1)
# 重试窗口内不再取数
self.assertEqual(etf_signal.gen_signals(rt), [])
self.assertEqual(self.client.get.call_count, 1)
self.assertIn(CODE, etf_signal._retry_at)
def test_retry_happens_after_window(self):
self.client.get.side_effect = boom
rt = runtime(self.config())
etf_signal.gen_signals(rt)
etf_signal._retry_at[CODE] = datetime.now() - timedelta(seconds=1)
self.client.get.side_effect = None
self.client.get.return_value = daily_response(raw_payload())
signals = etf_signal.gen_signals(rt)
self.assertEqual([item.code for item in signals], [CODE])
self.assertEqual(self.client.get.call_count, 2)
def test_new_trading_day_clears_cache(self):
rt = runtime(self.config())
etf_signal.gen_signals(rt)
# 模拟隔日:昨日的取数记录不应继续复用。
etf_signal._fetched[CODE] = LAST_DAY - timedelta(days=1)
etf_signal._daily_cache.clear()
etf_signal.gen_signals(rt)
self.assertEqual(etf_signal._fetched[CODE], datetime.now().date())
self.assertEqual(self.client.get.call_count, 2)
def test_endpoint_falls_back_to_default_without_api_host(self):
etf_signal.gen_signals(runtime(self.config(), host=""))
self.assertEqual(self.client.get.call_args.args[0], etf_signal.DAILY_URL)
def test_endpoint_uses_global_api_host(self):
etf_signal.gen_signals(runtime(self.config(), host="http://api.test/"))
self.assertEqual(self.client.get.call_args.args[0], "http://api.test/etf/daily")
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,358 @@
"""ETF 开仓与持仓:白名单/入场门槛/反弹确认、主出口、副出口、百分比补仓。"""
from datetime import date, datetime, timedelta
import unittest
from unittest.mock import Mock, patch
from config import AccountConfig, EtfConfig, EtfDefaults, EtfSymbolConfig, GlobalConfig
from libs.order import OrderBook
from libs.runtime import Runtime
from libs.signal import SignalItem
from libs import watch
from libs.watch import DipWatch
from sdk import Assets, PositionItem, Tick
from strategy.etf import open as etf_open
from strategy.etf import positions as etf_positions
CODE = "510300.SH"
OTHER = "159915.SZ"
# 固定"当前时刻",与 tick 的时间戳保持同一交易日且不过期。
NOW = datetime(2026, 9, 16, 10, 0, 0)
class FrozenDateTime(datetime):
"""冻结 ``datetime.now()``,其余行为与标准库一致。"""
@classmethod
def now(cls, tz=None):
return NOW if tz is None else NOW.astimezone(tz)
def stamp(now: datetime) -> str:
return now.strftime("%Y%m%d %H:%M:%S")
def tick(price: float, now: datetime | None = None) -> Tick:
now = now or NOW
return Tick(last_price=price, last_close=price, raw={"timetag": stamp(now)})
def symbol(**overrides) -> EtfSymbolConfig:
base = dict(is_t0=False, buy_shares=1000, atr_multiplier=1.0, inner_step=0.7)
base.update(overrides)
return EtfSymbolConfig(**base)
def etf_config(**symbol_overrides) -> EtfConfig:
return EtfConfig(
defaults=EtfDefaults(),
symbols={CODE: symbol(**symbol_overrides)},
)
def position(volume: int, cost: float, can_use: int | None = None, name: str = "") -> PositionItem:
return PositionItem(
stock_code=CODE,
stock_name=name,
volume=volume,
open_price=cost,
can_use_volume=volume if can_use is None else can_use,
yesterday_volume=volume if can_use is None else can_use,
last_price=cost,
)
def signal(entry: float = 9.65, price: float = 10.0, code: str = CODE) -> SignalItem:
return SignalItem(
signal_key="etf",
code=code,
last_close=price,
tech_indicator={
"etf_entry": entry,
"etf_price": price,
"etf_grid": 1.0,
"etf_add_price": price * 0.97,
"etf_ma60": 10.0,
},
)
class ETFTradeBase(unittest.TestCase):
def setUp(self):
self.client = Mock()
self.client.passorder.return_value = {"status": "success"}
# Runtime.__post_init__ 会拉一次服务端初始化数据,测试里不发真实请求。
patch("libs.runtime.get_json", side_effect=OSError("offline")).start()
self.run = Runtime(
client=self.client,
global_cfg=GlobalConfig(api_host="http://api.test"),
account_cfg=AccountConfig(account_id="acct", strategy="etf", min_cash_ratio=0.0),
etf_cfg=etf_config(),
orders=OrderBook(),
open_watch=DipWatch(expire_seconds=600, rebound_threshold=0.5),
add_watch=DipWatch(expire_seconds=600, rebound_threshold=0.5),
)
for module in (etf_open, etf_positions, watch):
patch.object(module, "trading_time", return_value=True, create=True).start()
patch.object(module, "datetime", FrozenDateTime).start()
patch.dict(etf_positions._progress, {}, clear=True).start()
patch.dict(etf_positions._trackers, {}, clear=True).start()
self.addCleanup(patch.stopall)
def last_order(self) -> dict:
self.assertTrue(self.client.passorder.called, "未提交任何委托")
return self.client.passorder.call_args.kwargs
class OpenSignalTests(ETFTradeBase):
def test_whitelist_outside_config_is_skipped(self):
item = signal(code=OTHER)
etf_open.open_signal(self.run, {OTHER: tick(9.4)}, [item])
self.client.passorder.assert_not_called()
def test_price_above_entry_no_observation_no_order(self):
etf_open.open_signal(self.run, {CODE: tick(9.9)}, [signal(entry=9.65)])
self.client.passorder.assert_not_called()
self.assertEqual(self.run.open_watch.data, {})
def test_seesaw_below_entry_requires_rebound_confirmation(self):
item = signal(entry=9.65)
etf_open.open_signal(self.run, {CODE: tick(9.4)}, [item]) # 进入入场区,记低点
etf_open.open_signal(self.run, {CODE: tick(9.39)}, [item]) # 刷新低点
etf_open.open_signal(self.run, {CODE: tick(9.40)}, [item]) # 反弹不足 0.5%
self.client.passorder.assert_not_called()
etf_open.open_signal(self.run, {CODE: tick(9.42)}, [item]) # (9.42-9.39)/9.39 = 0.32% 仍不足
self.client.passorder.assert_not_called()
def test_rebound_places_base_limit_order_at_anchor(self):
item = signal(entry=9.65)
etf_open.open_signal(self.run, {CODE: tick(9.40)}, [item])
etf_open.open_signal(self.run, {CODE: tick(9.45)}, [item]) # 反弹 0.53% 确认
request = self.last_order()
self.assertEqual(request["op_type"], 23)
self.assertEqual(request["volume"], 1000)
self.assertEqual(request["price"], 9.45)
self.assertEqual(request["pr_type"], 11)
self.assertEqual(request["strategy_name"], "etf")
self.assertEqual(self.run.open_watch.data, {})
def test_leaving_entry_band_forgets_the_watch(self):
item = signal(entry=9.65)
etf_open.open_signal(self.run, {CODE: tick(9.4)}, [item])
etf_open.open_signal(self.run, {CODE: tick(9.8)}, [item])
self.assertEqual(self.run.open_watch.data, {})
def test_missing_entry_indicator_is_skipped(self):
item = signal()
item.tech_indicator.clear()
etf_open.open_signal(self.run, {CODE: tick(9.4)}, [item])
self.client.passorder.assert_not_called()
def test_stale_tick_is_skipped(self):
old = datetime(2026, 9, 16, 9, 50, 0)
etf_open.open_signal(self.run, {CODE: tick(9.4, now=old)}, [signal(entry=9.65)])
self.client.passorder.assert_not_called()
def test_previous_day_tick_is_skipped(self):
yesterday = datetime(2026, 9, 15, 14, 0, 0)
etf_open.open_signal(self.run, {CODE: tick(9.4, now=yesterday)}, [signal(entry=9.65)])
self.client.passorder.assert_not_called()
def test_insufficient_budget_cancels_the_anchor(self):
self.run.client.assets.return_value = Assets(total=1000.0, available=100.0)
item = signal(entry=9.65)
etf_open.open_signal(self.run, {CODE: tick(9.40)}, [item])
etf_open.open_signal(self.run, {CODE: tick(9.45)}, [item])
self.client.passorder.assert_not_called()
self.assertEqual(self.run.open_watch.data, {})
class MainExitTests(ETFTradeBase):
def test_profit_above_target_clears_the_whole_grid(self):
etf_positions.manage_positions(
self.run, {CODE: tick(10.2)}, [position(2000, 10.0)], True, 100000.0
)
request = self.last_order()
self.assertEqual(request["op_type"], 24)
self.assertEqual(request["volume"], 2000)
self.assertEqual(request["strategy_name"], "etf")
self.assertEqual(request["price"], 10.2)
self.assertEqual(request["pr_type"], 11)
def test_profit_below_target_does_not_sell(self):
etf_positions.manage_positions(
self.run, {CODE: tick(10.05)}, [position(2000, 10.0)], True, 100000.0
)
self.client.passorder.assert_not_called()
def test_t_plus_1_position_bought_today_is_not_sellable(self):
held = position(1000, 10.0, can_use=0)
held.yesterday_volume = 0
etf_positions.manage_positions(
self.run, {CODE: tick(10.5)}, [held], True, 100000.0
)
self.client.passorder.assert_not_called()
def test_t0_symbol_sells_on_the_same_day(self):
self.run.etf_cfg = etf_config(is_t0=True)
held = position(1000, 10.0, can_use=1000)
held.yesterday_volume = 0
etf_positions.manage_positions(
self.run, {CODE: tick(10.5)}, [held], True, 100000.0
)
self.assertEqual(self.last_order()["volume"], 1000)
class LevelExitTests(ETFTradeBase):
"""副出口:主出口在盈亏率 ≥1% 时会先吃掉整仓,因此这里直接喂盈亏率验证峰值回撤。
inner_step = 0.7、inner_grids = 2只有峰值抬到第 2 格后的回撤才允许卖出。
"""
def observe(self, series, cost: float = 11.6):
held = position(1000, cost)
symbol = self.run.etf_cfg.symbols[CODE]
level = etf_positions.position_level(self.run, held)
decisions = []
for pnl_rate in series:
price = cost * (1 + pnl_rate / 100)
decisions.append(
etf_positions.handle_level_exit(
self.run, symbol, held, tick(price), pnl_rate, level
)
)
return decisions
def test_peak_retreat_sells_only_that_level(self):
first, second, third = self.observe([0.5, 1.5, 1.2])
self.assertFalse(first.submitted) # 首次观察,只建基准
self.assertFalse(second.submitted) # 峰值抬到第 2 格
self.assertTrue(third.submitted) # 回撤到第 1 格
request = self.last_order()
self.assertEqual(request["op_type"], 24)
self.assertEqual(request["volume"], 1000)
self.assertEqual(request["strategy_name"], "etf")
def test_retreat_below_inner_grids_is_held(self):
first, second = self.observe([0.5, 0.1])
self.assertFalse(first.submitted)
self.assertFalse(second.submitted) # 峰值只有 0 格
self.client.passorder.assert_not_called()
def test_peak_is_kept_when_the_order_is_rejected(self):
self.client.passorder.return_value = {"status": "rejected"}
self.run.orders.place = Mock(return_value=False)
first, second, third = self.observe([0.5, 1.5, 1.2])
self.assertFalse(third.submitted)
# 下单失败必须保留峰值:下一轮同样能再次触发。
fourth = self.observe([1.2])[0]
self.assertFalse(fourth.submitted)
def test_manage_positions_runs_the_secondary_exit(self):
"""成本 11.6、现价 11.7/11.65 的盈亏率都低于 1%,主出口不参与。"""
held = position(1000, 11.6)
etf_positions.manage_positions(self.run, {CODE: tick(11.7)}, [held], True, 0.0)
etf_positions.manage_positions(self.run, {CODE: tick(11.65)}, [held], True, 0.0)
self.client.passorder.assert_not_called() # 峰值未达 2 格
class AddTests(ETFTradeBase):
"""补仓规则单测:直接调 handle_add避免其它出口的委托锁干扰。"""
def held(self, volume: int = 3000, cost: float = 10.0) -> PositionItem:
return position(volume, cost)
def add(self, price: float, volume: int = 3000, cost: float = 10.0,
budget: float = 100000.0, level: int | None = None,
symbol_overrides: dict | None = None, first_low: float | None = None):
if symbol_overrides:
self.run.etf_cfg = etf_config(**symbol_overrides)
held = self.held(volume, cost)
if first_low is not None:
# 先造出一个观察低点,再由本次调用验证反弹确认。
self.run.add_watch.triggered("补仓", CODE, first_low)
return etf_positions.handle_add(
self.run,
self.run.etf_cfg.symbols[CODE],
held,
tick(price),
price,
budget,
level if level is not None else etf_positions.position_level(self.run, held),
)
def test_add_requires_add_pct_drop(self):
decision = self.add(9.95)
self.assertFalse(decision.submitted)
self.client.passorder.assert_not_called()
self.assertEqual(self.run.add_watch.data, {})
def test_add_waits_for_rebound_before_buying(self):
# 跌幅 4% ≥ add_pct 3%,但还没反弹确认:只观察,不下单。
first = self.add(9.60)
self.assertFalse(first.submitted)
self.assertIn(CODE, self.run.add_watch.data)
# 从观察低点 9.60 反弹 0.63%:确认后按现价买一档。
second = self.add(9.66)
self.assertTrue(second.submitted)
request = self.last_order()
self.assertEqual(request["op_type"], 23)
self.assertEqual(request["volume"], 1000)
self.assertEqual(request["price"], 9.66)
self.assertEqual(request["strategy_name"], "etf")
self.assertEqual(self.run.add_watch.data, {})
def test_add_stops_at_max_adds(self):
decision = self.add(9.60, volume=10000)
self.assertFalse(decision.submitted)
self.assertIn("", decision.message)
def test_add_respects_max_shares(self):
decision = self.add(9.60, volume=3000, symbol_overrides={"max_shares": 3000})
self.assertFalse(decision.submitted)
self.client.passorder.assert_not_called()
def test_add_needs_budget(self):
# 跌幅 4%、反弹 0.63% 都满足,但预算为 0不消耗观察状态也不下单。
first = self.add(9.60, budget=0.0)
self.assertFalse(first.submitted)
second = self.add(9.66, budget=0.0, first_low=9.60)
self.assertFalse(second.submitted)
self.assertIn(CODE, self.run.add_watch.data)
# 资金到位后同一个观察低点仍可确认。
third = self.add(9.66, budget=100000.0, first_low=None)
self.assertTrue(third.submitted)
def test_add_uses_broker_cost_as_previous_level(self):
# 上一档 = 券商成本 9.5:跌到 9.16 是 3.58% ≥ 3%,反弹到 9.21 确认。
decision = self.add(9.21, cost=9.5, first_low=9.16)
self.assertTrue(decision.submitted)
self.assertEqual(self.last_order()["price"], 9.21)
def test_add_blocked_when_market_disallows(self):
held = self.held()
etf_positions.manage_positions(self.run, {CODE: tick(9.6)}, [held], False, 100000.0)
etf_positions.manage_positions(self.run, {CODE: tick(9.66)}, [held], False, 100000.0)
self.client.passorder.assert_not_called()
def test_add_blocked_by_in_flight_buy(self):
self.run.orders.busy_cache.set("BUY-" + CODE, True, timeout=180)
self.add(9.60)
self.add(9.66)
self.client.passorder.assert_not_called()
class PositionLevelTests(ETFTradeBase):
def test_level_is_derived_from_volume(self):
for volume, expected in ((1000, 1), (2000, 2), (3500, 4), (10000, 10)):
with self.subTest(volume=volume):
self.assertEqual(
etf_positions.position_level(self.run, position(volume, 10.0)), expected
)
def test_unknown_symbol_yields_baseline_level(self):
self.assertEqual(etf_positions.position_level(self.run, position(1000, 10.0)), 1)
if __name__ == "__main__":
unittest.main()

View File

@@ -3,7 +3,7 @@ import io
import logging
import unittest
from concurrent.futures import Future
from contextlib import ExitStack, redirect_stdout
from contextlib import ExitStack, redirect_stderr, redirect_stdout
from types import SimpleNamespace
from unittest.mock import Mock, patch
@@ -77,7 +77,7 @@ class TrendCollectorTests(unittest.TestCase):
self.assertEqual(payload['deals'][0]['volume'], 100)
def test_main_registers_five_minute_collector_job(self):
for strategy in ('trend', 'zt'):
for strategy in ('trend', 'zt', 'etf'):
with self.subTest(strategy=strategy), ExitStack() as stack:
scheduler = Mock(running=True)
stack.enter_context(patch.object(self.app, 'BackgroundScheduler', return_value=scheduler))
@@ -101,6 +101,63 @@ class TrendCollectorTests(unittest.TestCase):
scheduler.start.assert_called_once()
scheduler.shutdown.assert_called_once_with(wait=True)
def test_main_rejects_unknown_strategy_before_starting_the_scheduler(self):
scheduler = Mock(running=True)
with patch.object(self.app, 'BackgroundScheduler', return_value=scheduler), \
patch.object(self.app, 'require_windows', return_value=True), \
patch.object(self.app, 'check_single_instance', return_value=True), \
patch.object(self.app, 'wait_for_qmt_api') as wait_api, \
patch.object(self.app.config, 'load'), \
patch.object(self.app.config, 'global_config', SimpleNamespace(api_host='unused')), \
patch.object(self.app.config, 'account_config', SimpleNamespace(strategy='bogus')), \
patch.object(self.app, 'wait_for_any_key'), \
patch('sys.stdin', io.StringIO('\n')), redirect_stderr(io.StringIO()) as err:
self.assertEqual(self.app.main(), 1)
self.assertIn('bogus', err.getvalue())
scheduler.start.assert_not_called()
wait_api.assert_not_called()
def test_main_reports_a_missing_global_config(self):
scheduler = Mock(running=True)
with patch.object(self.app, 'BackgroundScheduler', return_value=scheduler), \
patch.object(self.app, 'require_windows', return_value=True), \
patch.object(self.app, 'check_single_instance', return_value=True), \
patch.object(self.app, 'wait_for_qmt_api') as wait_api, \
patch.object(self.app.config, 'load'), \
patch.object(self.app.config, 'global_config', None), \
patch.object(self.app.config, 'account_config', None), \
patch.object(self.app, 'wait_for_any_key'), \
patch('sys.stdin', io.StringIO('\n')), redirect_stderr(io.StringIO()) as err:
self.assertEqual(self.app.main(), 1)
self.assertIn('config.load', err.getvalue())
scheduler.add_job.assert_not_called()
wait_api.assert_not_called()
class ETFEtfConfigSummaryTests(unittest.TestCase):
"""启动日志里的 ETF 配置概览:缺文件必须说清楚,且不抛异常。"""
def summary(self, etf_cfg):
app = importlib.import_module('main')
with patch.object(app.config, 'etf_config', etf_cfg, create=True):
return app.describe_etf_config()
def test_missing_file_is_described_not_raised(self):
self.assertIn('_etf.yaml', self.summary(None))
def test_symbols_and_codes_are_listed_in_config_order(self):
from config import EtfConfig, EtfDefaults, EtfSymbolConfig
cfg = EtfConfig(
defaults=EtfDefaults(),
symbols={
'159915.SZ': EtfSymbolConfig(is_t0=True, buy_shares=1000, atr_multiplier=1.0, inner_step=0.7),
'510300.SH': EtfSymbolConfig(is_t0=False, buy_shares=1000, atr_multiplier=1.0, inner_step=0.7),
},
)
summary = self.summary(cfg)
self.assertIn('2 只', summary)
self.assertLess(summary.index('159915.SZ'), summary.index('510300.SH'))
if __name__ == '__main__':
unittest.main()