133 lines
6.7 KiB
Python
133 lines
6.7 KiB
Python
import tempfile
|
|
import unittest
|
|
from contextlib import closing
|
|
from dataclasses import replace
|
|
from datetime import datetime, timedelta
|
|
from pathlib import Path
|
|
from types import SimpleNamespace as NS
|
|
from unittest.mock import Mock, patch
|
|
|
|
from libs.grid_take_profit import GridState
|
|
from libs.order import OrderBook
|
|
from libs.snapshot import get_collector_snapshot
|
|
from libs.state import State
|
|
from sdk import Assets, DealItem, OrderItem, PositionItem
|
|
from strategy.zt import boot
|
|
from strategy.zt.profit import ZTProfitTracker
|
|
|
|
|
|
class ZTAuditFixTests(unittest.TestCase):
|
|
def setUp(self):
|
|
tmp = tempfile.TemporaryDirectory()
|
|
self.addCleanup(tmp.cleanup)
|
|
self.store = State(Path(tmp.name) / 'state.db')
|
|
self.code = '600000.SH'
|
|
|
|
def deal(self, code, identity):
|
|
return DealItem(stock_code=code, order_sys_id=identity,
|
|
remark=f'zt-base-{identity}|zt', offset_flag=48,
|
|
volume=100, price=10, trade_amount=1000)
|
|
|
|
def position(self, code):
|
|
return PositionItem(stock_code=code, volume=100, open_price=10)
|
|
|
|
def test_zt_cancels_only_owned_orders_and_tracks_all(self):
|
|
stamp = datetime.now() - timedelta(minutes=2)
|
|
orders = [OrderItem(stock_code=f'60000{i}.SH', order_sys_id=str(i), remark=remark,
|
|
order_status=50, offset_flag=23,
|
|
insert_date=stamp.strftime('%Y%m%d'), insert_time=stamp.strftime('%H%M%S'))
|
|
for i, remark in enumerate(['zt-base-own|zt', 'zt-SELL-own|zt',
|
|
'zt-added-own|zt', 'IPO-new|ipo', '', 'TREN-BUY-other'])]
|
|
client = Mock()
|
|
book = OrderBook()
|
|
book.refresh(client, orders, cancel_prefix='zt-')
|
|
self.assertEqual([c.args[0] for c in client.cancel_by_id.call_args_list], ['0', '1', '2'])
|
|
self.assertEqual(book.data, orders)
|
|
self.assertTrue(all(book.busy(o.stock_code, 'BUY') for o in orders))
|
|
client.reset_mock()
|
|
book.refresh(client, orders)
|
|
self.assertEqual([c.args[0] for c in client.cancel_by_id.call_args_list], ['0', '1', '2', '4', '5'])
|
|
|
|
def test_invalid_trade_persists_blocks_only_its_stock_and_recovers(self):
|
|
good = '600001.SH'
|
|
bad = replace(self.deal(self.code, 'bad'), price=float('nan'))
|
|
valid = self.deal(good, 'good')
|
|
positions = [self.position(c) for c in (self.code, good)]
|
|
self.store.sync_account(positions, [bad, valid])
|
|
self.assertEqual(self.store.blocked_codes, {self.code})
|
|
self.assertEqual(self.store.state[good]['base_qty'], 100)
|
|
with closing(self.store._connect()) as db:
|
|
payload = db.execute('SELECT payload FROM zt_rejected_deals').fetchone()[0]
|
|
self.assertIn('bad', payload)
|
|
self.assertIn('NaN', payload)
|
|
self.store = State(self.store.path)
|
|
self.store.sync_account(positions, [valid])
|
|
self.assertEqual(self.store.blocked_codes, {self.code})
|
|
self.store.sync_account(positions, [replace(bad, price=10), valid])
|
|
self.assertEqual(self.store.blocked_codes, set())
|
|
self.assertEqual(self.store.state[self.code]['base_qty'], 100)
|
|
|
|
def test_invalid_fields_do_not_block_other_stocks(self):
|
|
for field, value in [('volume', 0), ('volume', 1.5), ('offset_flag', 99),
|
|
('trade_amount', float('inf')), ('price', -1)]:
|
|
with self.subTest(field=field):
|
|
bad = replace(self.deal(self.code, f'bad-{field}'), **{field: value})
|
|
good = self.deal('600001.SH', 'good')
|
|
self.store.sync_account([self.position('600001.SH')], [bad, good])
|
|
self.assertEqual(self.store.state['600001.SH']['base_qty'], 100)
|
|
self.assertIn(self.code, self.store.blocked_codes)
|
|
|
|
def test_bad_stock_does_not_archive_other_trades_until_corrected(self):
|
|
first = self.deal(self.code, 'first')
|
|
bad = replace(self.deal(self.code, 'bad'), price=float('nan'))
|
|
position = replace(self.position(self.code), volume=200)
|
|
self.store.sync_account([position], [first, bad])
|
|
self.assertEqual(self.store.deals['first']['is_arch'], 0)
|
|
self.assertNotIn(self.code, self.store.state)
|
|
self.store.sync_account([position], [replace(bad, price=10)])
|
|
self.assertEqual(self.store.state[self.code]['base_qty'], 200)
|
|
self.assertEqual(self.store.blocked_codes, set())
|
|
|
|
def test_profit_basis_changes_reset_peak_but_partial_sell_does_not(self):
|
|
tracker = ZTProfitTracker()
|
|
position = self.position(self.code)
|
|
row = dict(base_qty=100, base_order_local_id='one', base_created_at='now', added_qty=0)
|
|
state = NS(blocked_codes=set(), get_by_code=lambda code: row)
|
|
tracker.sync_positions([position], state)
|
|
self.assertEqual(tracker.observe(self.code, 20).state, GridState.ARMED)
|
|
row['base_qty'] = 50
|
|
tracker.sync_positions([replace(position, volume=50)], state)
|
|
self.assertEqual(tracker.observe(self.code, 19).state, GridState.RETREAT)
|
|
for update, cost in [({'added_qty': 100, 'added_price': 11}, 10),
|
|
({'added_qty': 0}, 10), ({}, 12),
|
|
({'base_order_local_id': 'reopened'}, 12)]:
|
|
row.update(update)
|
|
tracker.sync_positions([replace(position, open_price=cost)], state)
|
|
self.assertEqual(tracker.observe(self.code, 10).state, GridState.ARMED)
|
|
tracker.sync_positions([], state)
|
|
tracker.sync_positions([position], state)
|
|
self.assertEqual(tracker.observe(self.code, 5).state, GridState.ARMED)
|
|
|
|
def test_run_once_updates_collector_before_market_fetch(self):
|
|
positions = [self.position(self.code)]
|
|
self.store.sync_account(positions, [], initialize=True)
|
|
client = Mock()
|
|
run = NS(client=client, account_cfg=NS(account_id='zt-test', min_cash_ratio=0.1),
|
|
orders=Mock(), profit_tracker=ZTProfitTracker())
|
|
client.deals.return_value = []
|
|
client.portfolio.return_value = NS(assets=Assets(20000, 10000),
|
|
positions={self.code: positions[0]}, orders=[])
|
|
client.full_tick.side_effect = RuntimeError('no market data')
|
|
with patch.object(boot, 'trading_time', return_value=True), \
|
|
patch.object(boot, 'market_allow_open', return_value=True):
|
|
boot.RunOnce(run, self.store, [])
|
|
snapshot = get_collector_snapshot()
|
|
self.assertEqual(snapshot[0], 'zt-test')
|
|
self.assertEqual(snapshot[1].total, 20000)
|
|
self.assertEqual(snapshot[2], positions)
|
|
run.orders.refresh.assert_called_once_with(client, [], cancel_prefix='zt-')
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|