fix bug
This commit is contained in:
@@ -52,6 +52,20 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertAlmostEqual(row['added_price'], 1620 / 150)
|
||||
self.assertEqual(row['added_order_local_id'], 'late')
|
||||
|
||||
def test_no_argument_archiving_recognizes_base_and_added_orders(self):
|
||||
self.insert_deal('zt-base-first', 100, 1000, '10:00:00', flag=23)
|
||||
self.insert_deal('zt-base-second', 100, 1200, '10:01:00', flag=48)
|
||||
self.insert_deal('zt-t-buy-first', 100, 900, '10:02:00', flag=23)
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['base_price']), (200, 11))
|
||||
self.assertEqual((row['added_qty'], row['added_price']), (100, 9))
|
||||
self.assertEqual(row['base_order_local_id'], 'zt-base-second')
|
||||
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
|
||||
saved = dict(row)
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.assertEqual(self.book.state['600000.SH'], saved)
|
||||
|
||||
def test_sell_added_then_clear_base(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.insert_deal('buy1', 40, 400, '09:59:00')
|
||||
@@ -88,8 +102,9 @@ class ArchivingTests(unittest.TestCase):
|
||||
def test_excess_sell_rolls_back(self):
|
||||
self.insert_deal('buy', 50, 500, '10:00:00')
|
||||
self.insert_deal('sell', 100, 1200, '10:01:00', flag=49)
|
||||
errors = self.book.archiving()
|
||||
self.assertIn('600000.SH', errors)
|
||||
with self.assertLogs(level='WARNING') as logs:
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.assertIn('600000.SH', '\n'.join(logs.output))
|
||||
restarted = State(self.book.path)
|
||||
self.assertEqual(restarted.state, {})
|
||||
self.assertTrue(all(deal['is_arch'] == 0 for deal in restarted.deals.values()))
|
||||
@@ -98,7 +113,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.insert_deal('buy', 100, 1000, '10:00:00', flag=23)
|
||||
self.insert_deal('sell', 50, 600, '10:01:00', flag=24)
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 50))
|
||||
self.assertEqual(self.book.deals['buy']['offset_flag'], 23)
|
||||
@@ -113,8 +128,9 @@ class ArchivingTests(unittest.TestCase):
|
||||
db.execute("""CREATE TRIGGER fail_archive BEFORE UPDATE OF is_arch ON deals
|
||||
WHEN OLD.order_sys_id = 'second'
|
||||
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
|
||||
errors = self.book.archiving()
|
||||
self.assertIn('600001.SH', errors)
|
||||
with self.assertLogs(level='WARNING') as logs:
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.assertIn('600001.SH', '\n'.join(logs.output))
|
||||
restarted = State(self.book.path)
|
||||
self.assertEqual(set(restarted.state), {'600000.SH'})
|
||||
self.assertEqual(restarted.state, self.book.state)
|
||||
@@ -134,17 +150,17 @@ class ArchivingTests(unittest.TestCase):
|
||||
offset_flag=48, volume=100, price=8, trade_amount=800,
|
||||
trade_date='20260908', trade_time='100000',
|
||||
)])
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 100))
|
||||
self.assertEqual(self.book.deals['first']['is_arch'], 1)
|
||||
restarted = State(self.book.path)
|
||||
self.assertEqual(restarted.archiving(), {})
|
||||
self.assertIsNone(restarted.archiving())
|
||||
self.assertEqual(restarted.state, self.book.state)
|
||||
self.insert_deal('new_buy', 100, 1000, '10:01:00')
|
||||
with closing(self.book._connect()) as db, db:
|
||||
db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'")
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 200, 'CUSTOM'))
|
||||
|
||||
@@ -153,8 +169,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
with closing(self.book._connect()) as db, db:
|
||||
db.execute('UPDATE deals SET is_arch = 1')
|
||||
self.insert_deal('new', 50, 500, '10:01:00')
|
||||
errors = self.book.archiving()
|
||||
self.assertEqual(errors, {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.assertEqual(self.book.state['600000.SH']['added_qty'], 50)
|
||||
self.assertEqual(self.book.deals['old']['is_arch'], 1)
|
||||
self.assertEqual(self.book.deals['new']['is_arch'], 1)
|
||||
@@ -166,7 +181,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(self.book.state, {})
|
||||
self.book = State(self.book.path)
|
||||
self.insert_deal('late', 100, 2000, '100100')
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual(row['added_qty'], 100)
|
||||
self.assertEqual(row['added_price'], 20)
|
||||
@@ -174,21 +189,23 @@ class ArchivingTests(unittest.TestCase):
|
||||
def test_archive_sell_before_syncing_empty_positions(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.insert_deal('sell', 100, 1000, '10:01:00', flag=49)
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.book.sync_state([])
|
||||
self.book = State(self.book.path)
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.assertEqual(self.book.state, {})
|
||||
self.assertEqual(self.book.deals['sell']['is_arch'], 1)
|
||||
|
||||
def test_bad_stock_does_not_block_good_stock_and_can_retry(self):
|
||||
self.insert_deal('bad_sell', 100, 1500, '10:02:00', flag=49)
|
||||
self.insert_deal('good_buy', 100, 1000, '10:00:00', code='600001.SH')
|
||||
self.assertIn('600000.SH', self.book.archiving())
|
||||
with self.assertLogs(level='WARNING') as logs:
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.assertIn('600000.SH', '\n'.join(logs.output))
|
||||
self.assertEqual(self.book.deals['good_buy']['is_arch'], 1)
|
||||
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 0)
|
||||
self.insert_deal('late_buy', 100, 1000, '10:01:00')
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.assertIsNone(self.book.archiving())
|
||||
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 1)
|
||||
self.assertNotIn('600000.SH', self.book.state)
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
from pathlib import Path
|
||||
|
||||
from libs.state import State
|
||||
@@ -49,6 +50,32 @@ class ZTStateTests(unittest.TestCase):
|
||||
self.assertEqual(self.state.state['600000.SH']['base_qty'], 100)
|
||||
self.assertEqual(self.state.deals['sell']['is_arch'], 0)
|
||||
|
||||
def test_failed_initialization_leaves_original_database_empty(self):
|
||||
invalid = self.position(100)
|
||||
invalid.open_price = float('inf')
|
||||
with self.assertRaises(ValueError):
|
||||
sync_account_state(self.state, [invalid], [self.deal('old', 100)], initialize=True)
|
||||
restarted = State(self.state.path)
|
||||
self.assertEqual((restarted.state, restarted.deals), ({}, {}))
|
||||
sync_account_state(restarted, [self.position(100)], [self.deal('old', 100)], initialize=True)
|
||||
self.assertEqual(restarted.deals['old']['is_arch'], 1)
|
||||
|
||||
def test_initialization_cannot_overwrite_existing_database(self):
|
||||
sync_account_state(self.state, [self.position(100)], [], initialize=True)
|
||||
with self.assertRaises(ValueError):
|
||||
sync_account_state(self.state, [], [], initialize=True)
|
||||
self.assertEqual(State(self.state.path).state['600000.SH']['base_qty'], 100)
|
||||
|
||||
def test_no_new_deals_still_retries_failed_archiving(self):
|
||||
sync_account_state(self.state, [self.position(100)], [], initialize=True)
|
||||
sell = self.deal('sell', 100, flag=24)
|
||||
with patch.object(self.state, 'archiving'):
|
||||
with self.assertRaises(ValueError):
|
||||
sync_account_state(self.state, [], [sell])
|
||||
sync_account_state(self.state, [], [sell])
|
||||
self.assertEqual(self.state.state, {})
|
||||
self.assertEqual(self.state.deals['sell']['is_arch'], 1)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
||||
@@ -7,6 +7,7 @@ from unittest.mock import Mock, patch
|
||||
|
||||
from config import AccountConfig
|
||||
from libs.grid_take_profit import GridState
|
||||
from libs.order import OrderBook
|
||||
from libs.state import State
|
||||
from sdk import Assets, DealItem, PositionItem, Tick
|
||||
from strategy.zt import boot
|
||||
@@ -24,7 +25,7 @@ class ZTTradingTests(unittest.TestCase):
|
||||
self.run = SimpleNamespace(account_cfg=self.cfg, orders=Mock(), client=Mock(),
|
||||
profit_tracker=Mock(), add_watch=Mock(), open_watch=Mock())
|
||||
self.run.orders.busy.return_value = False
|
||||
self.run.orders.new_order_id.side_effect = lambda kind: f'zt-{kind}-order'
|
||||
self.run.orders.new_order_id.side_effect = lambda prefix, kind: f'{prefix}-{kind}-order'
|
||||
self.run.profit_tracker.observe.return_value.state = GridState.RETREAT
|
||||
self.run.add_watch.triggered.return_value = True
|
||||
self.run.open_watch.triggered.return_value = True
|
||||
@@ -131,6 +132,21 @@ class ZTTradingTests(unittest.TestCase):
|
||||
[SimpleNamespace(code='688001.SH')], 2000)
|
||||
self.run.orders.place.assert_not_called()
|
||||
|
||||
def test_real_order_id_is_recognized_by_state_sync(self):
|
||||
orders = OrderBook('zt')
|
||||
self.run.orders.new_order_id.side_effect = orders.new_order_id
|
||||
with patch('strategy.zt.open.datetime') as clock:
|
||||
clock.now.return_value = datetime(2026, 9, 9, 10)
|
||||
open_signal(self.run, {self.code: Tick(last_price=10)},
|
||||
[SimpleNamespace(code=self.code)], 2000)
|
||||
request = self.run.orders.place.call_args.args[1]
|
||||
self.assertTrue(request.order_id.startswith('zt-base-'))
|
||||
deal = self.fill('base', 'b1', 100)
|
||||
deal.remark = request.order_id + '|zt'
|
||||
self.store.sync_state([])
|
||||
boot.sync_account_state(self.store, [], [deal])
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 100)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user