87 lines
4.1 KiB
Python
87 lines
4.1 KiB
Python
import tempfile
|
|
import unittest
|
|
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, 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 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_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_state(positions)
|
|
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()
|