add etf
This commit is contained in:
329
py-client/tests/test_etf.py
Normal file
329
py-client/tests/test_etf.py
Normal 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()
|
||||
Reference in New Issue
Block a user