This commit is contained in:
2026-09-17 00:56:05 +08:00
parent 33693adb66
commit c7d39938a2
14 changed files with 1005 additions and 0 deletions

329
py-client/tests/test_etf.py Normal file
View File

@@ -0,0 +1,329 @@
"""ETF 离线回归:指标、真实防飞刀/网格算法、限仓、回报和持久化。"""
from datetime import date, datetime, timedelta
import httpx
from pathlib import Path
import tempfile
import unittest
from unittest.mock import Mock
from sdk import Assets, OrderItem, Portfolio, PositionItem, Tick
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
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)
def tick(price, now=NOW):
return Tick(price, raw={'timetag': now.strftime('%Y%m%d %H:%M:%S')})
def position(volume=0, cost=0, available=None):
return PositionItem(stock_code=CODE, volume=volume, open_price=cost,
can_use_volume=volume if available is None else available)
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)
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 buy(self):
self.run_price(9.4)
self.run_price(9.46)
self.client.passorder.assert_called_once()
def report(self, status=56, filled=100, side=23, price=9.46):
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)
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)
self.client.passorder.assert_not_called()
self.run_price(9.36)
request = self.client.passorder.call_args.kwargs
self.assertEqual((request['volume'], request['price'], request['pr_type']), (100, 9.36, 11))
self.assertEqual(request['strategy_name'], 'etf')
self.assertTrue(self.store.get(CODE).pending)
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'])
raise TimeoutError('unknown result')
self.client.passorder.side_effect = submit
self.run_price(9.4)
with self.assertLogs(level='ERROR'):
self.run_price(9.46)
self.engine = Engine(self.client, self.cfg, Store(self.path, 'test'), 0.1)
with self.assertLogs(level='WARNING'):
self.run_price(9.3, 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_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_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))
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.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)
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)
self.client.passorder.assert_called_once()
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.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):
pos = position()
pos.on_road_volume = 100
self.run_price(9.4, pos)
self.run_price(9.46, pos)
unknown = OrderItem(stock_code=CODE, order_status=255)
self.run_price(9.4, orders=[unknown])
self.run_price(9.46, 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.client.passorder.assert_not_called()
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_corrupt_state_does_not_silently_start_empty(self):
self.path.write_text('{', encoding='utf-8')
with self.assertRaises(ValueError):
Store(self.path, 'test')
class IndicatorTests(unittest.TestCase):
def bars(self):
days = []
day = date(2026, 9, 15)
while len(days) < 80:
if day.weekday() < 5:
days.append(day)
day -= timedelta(days=1)
return [dict(date=d.strftime('%Y%m%d'), high=11, low=9, close=10) for d in reversed(days)]
def test_known_constant_series_and_exclusion_of_unfinished_day(self):
cfg = ETFConfig(codes=(CODE,))
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))
def test_atr_accounts_for_gap_and_uses_wilder_smoothing(self):
rows = self.bars()
rows[-1].update(high=14, low=12, close=13)
ind = calculate(rows, NOW.date(), ETFConfig(codes=(CODE,)))
self.assertAlmostEqual(ind.atr, (2 * 13 + 4) / 14)
self.assertAlmostEqual(ind.ma60, 10.05)
self.assertGreater(ind.upper, ind.middle)
def test_bad_or_insufficient_history_is_rejected(self):
cfg = ETFConfig(codes=(CODE,))
for rows in (self.bars()[:59], self.bars() + [self.bars()[-1]],
self.bars()[:-1] + [dict(self.bars()[-1], close=float('nan'))]):
with self.assertRaises(ValueError):
calculate(rows, NOW.date(), cfg)
def test_grid_floor_and_tick_rounding(self):
rows = [dict(row, high=10.001, low=9.999) for row in self.bars()]
ind = calculate(rows, NOW.date(), ETFConfig(codes=(CODE,), min_grid_pct=0.501))
self.assertEqual(ind.grid, 0.051)
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': ()}):
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))
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 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()]})
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_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)])
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)
def test_http_error_and_invalid_json_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:
with self.assertRaises((httpx.HTTPStatusError, ValueError)):
daily_bars(client, CODE, NOW.date())
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}):
with self.subTest(payload=payload), self.assertRaises(ValueError):
parse_daily(payload, CODE, NOW.date())
def test_wrong_symbol_duplicate_dates_and_invalid_ohlc_are_rejected(self):
for rows in ([self.row(ts_code=OTHER)], [self.row(), self.row()],
[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())
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())
ind = calculate(rows, NOW.date(), ETFConfig(codes=(CODE,)))
self.assertEqual((ind.ma60, ind.atr, ind.grid), (10, 2, 2))
if __name__ == '__main__':
unittest.main()