220 lines
11 KiB
Python
220 lines
11 KiB
Python
import sqlite3
|
|
import tempfile
|
|
import unittest
|
|
from contextlib import closing
|
|
from pathlib import Path
|
|
from datetime import datetime
|
|
from unittest.mock import patch
|
|
|
|
from libs.orderbook import OrderBook
|
|
from sdk import DealItem, PositionItem
|
|
from strategy.zt.state import DONE, READY, SOLD, TState
|
|
|
|
|
|
class OrderBookTests(unittest.TestCase):
|
|
def setUp(self):
|
|
clock = patch('strategy.zt.state.datetime')
|
|
self.clock = clock.start()
|
|
self.addCleanup(clock.stop)
|
|
self.clock.now.return_value = datetime(2026, 9, 1)
|
|
self.tmp = tempfile.TemporaryDirectory()
|
|
self.addCleanup(self.tmp.cleanup)
|
|
self.path = Path(self.tmp.name) / 'state.db'
|
|
|
|
def deal(self, kind, sys_order_id, qty, price, date='2026-09-01'):
|
|
prefix = {'base': 'zt-base-', 'sell': 'zt-t-sell-', 'buy': 'zt-t-buy-'}[kind]
|
|
return DealItem.from_trade_detail({
|
|
'm_strOrderSysID': sys_order_id,
|
|
'm_strInstrumentID': '600000', 'm_strExchangeID': 'SH',
|
|
'm_strInstrumentName': 'Test stock',
|
|
'm_nOffsetFlag': '24' if kind == 'sell' else '23',
|
|
'm_nOrderStatus': '56', 'm_nVolumeTotal': '0',
|
|
'm_nVolumeTraded': str(qty), 'm_nOrderTime': '100000',
|
|
'm_strInsertDate': date, 'm_strInsertTime': '10:00:00',
|
|
'm_strRemark': prefix + 'order1|zt',
|
|
'm_dPrice': str(price + 1), 'm_dTradePrice': str(price),
|
|
'm_dTradeAmount': str(qty * price),
|
|
})
|
|
|
|
def test_partial_fills_restart_dedup_and_daily_cycle(self):
|
|
state = TState(self.path)
|
|
state.reconcile([PositionItem(stock_code='600000.SH', volume=200, open_price=10)], [])
|
|
self.assertEqual(state.deals, [])
|
|
first = self.deal('sell', 'd1', 40, 12)
|
|
second = self.deal('sell', 'd2', 60, 13)
|
|
state.reconcile([], [first])
|
|
state = TState(self.path)
|
|
self.assertEqual(state.items['600000.SH'].phase, SOLD)
|
|
self.assertEqual(state.items['600000.SH'].sell_qty, 40)
|
|
state.reconcile([], [first, first, second])
|
|
self.assertEqual(len(state.deals), 2)
|
|
self.assertAlmostEqual(state.items['600000.SH'].sell_price, 12.6)
|
|
self.clock.now.return_value = datetime.fromisoformat('2026-09-02')
|
|
state.reconcile([], [first, second])
|
|
self.assertEqual(len(state.deals), 2)
|
|
self.assertEqual(state.items['600000.SH'].phase, SOLD)
|
|
b1 = self.deal('buy', 'd3', 40, 11, '2026-09-02')
|
|
b2 = self.deal('buy', 'd4', 60, 10, '2026-09-02')
|
|
state.reconcile([], [b1])
|
|
self.assertEqual(state.items['600000.SH'].phase, SOLD)
|
|
state.reconcile([], [b1, b2])
|
|
self.assertEqual(TState(self.path).items['600000.SH'].phase, DONE)
|
|
self.assertAlmostEqual(state.items['600000.SH'].buy_cost, 10.4)
|
|
self.clock.now.return_value = datetime.fromisoformat('2026-09-03')
|
|
state.reconcile([], [])
|
|
item = TState(self.path).items['600000.SH']
|
|
self.assertEqual((item.phase, item.base_qty, item.base_cost, item.sell_qty), (READY, 200, 10, 0))
|
|
|
|
def test_json_is_never_read(self):
|
|
legacy = self.path.with_suffix('.json')
|
|
legacy.write_text('invalid JSON', encoding='utf-8')
|
|
book = OrderBook(self.path)
|
|
self.assertIsNone(book.load())
|
|
self.assertEqual((book.positions, book.deals, book.deals_sys_ids), ({}, {}, set()))
|
|
self.assertEqual(legacy.read_text(encoding='utf-8'), 'invalid JSON')
|
|
|
|
def test_sync_deals_deduplicates_batch_and_restart(self):
|
|
book = OrderBook(self.path)
|
|
self.assertEqual((book.positions, book.deals, book.deals_sys_ids), ({}, {}, set()))
|
|
first = self.deal('base', 'd1', 40, 10, '20260901')
|
|
second = self.deal('base', 'd2', 60, 12)
|
|
with patch.object(book, '_insert_deals', wraps=book._insert_deals) as insert:
|
|
book.sync_deals([first, first, second])
|
|
insert.assert_called_once()
|
|
self.assertEqual(len(insert.call_args.args[1]), 2)
|
|
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
|
|
self.assertEqual(book.deals['d1']['insert_date'], '2026-09-01')
|
|
self.assertEqual(book.deals['d2']['traded_volume'], 60)
|
|
book = OrderBook(self.path)
|
|
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
|
|
with patch.object(book, '_insert_deals') as insert:
|
|
book.sync_deals([first, second])
|
|
book.sync_deals([])
|
|
insert.assert_not_called()
|
|
self.assertEqual(len(book.deals), 2)
|
|
|
|
def test_sync_deals_failure_rolls_back_entire_batch_and_cache(self):
|
|
book = OrderBook(self.path)
|
|
first = self.deal('base', 'd1', 100, 10)
|
|
invalid = self.deal('base', 'd2', 100, 10)
|
|
invalid.side = 'INVALID'
|
|
with self.assertRaises(sqlite3.IntegrityError):
|
|
book.sync_deals([first, invalid])
|
|
self.assertEqual(book.deals, {})
|
|
self.assertEqual(book.deals_sys_ids, set())
|
|
self.assertEqual(OrderBook(self.path).deals, {})
|
|
invalid.side = 'BUY'
|
|
book.sync_deals([first, invalid])
|
|
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
|
|
|
|
def test_load_refreshes_all_caches(self):
|
|
book = OrderBook(self.path)
|
|
writer = OrderBook(self.path)
|
|
writer.save({'600000.SH': {'code': '600000.SH', 'status': READY}}, [])
|
|
writer.sync_deals([self.deal('base', 'd1', 100, 10)])
|
|
book.load()
|
|
self.assertEqual(book.positions['600000.SH']['status'], READY)
|
|
self.assertEqual(book.deals_sys_ids, {'d1'})
|
|
self.assertEqual(book.deals['d1']['local_order_id'], 'zt-base-order1')
|
|
|
|
def test_first_start_after_full_sale_keeps_buyback_quantity(self):
|
|
state = TState(self.path)
|
|
state.reconcile([], [self.deal('sell', 'd1', 100, 12)])
|
|
item = state.items['600000.SH']
|
|
self.assertEqual((item.base_qty, item.sell_qty, item.phase), (100, 100, SOLD))
|
|
state.reconcile([], [self.deal('buy', 'd2', 100, 11)])
|
|
self.assertEqual(TState(self.path).items['600000.SH'].phase, DONE)
|
|
|
|
def test_failed_insert_rolls_back_memory_and_database(self):
|
|
state = TState(self.path)
|
|
with closing(sqlite3.connect(self.path)) as db:
|
|
db.execute("""CREATE TRIGGER fail_insert BEFORE INSERT ON deals
|
|
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
|
|
fill = self.deal('base', 'd1', 100, 10)
|
|
with self.assertRaises(sqlite3.IntegrityError):
|
|
state.reconcile([], [fill])
|
|
self.assertFalse(state.items)
|
|
self.assertFalse(state.deals)
|
|
self.assertFalse(TState(self.path).items)
|
|
with closing(sqlite3.connect(self.path)) as db:
|
|
db.execute('DROP TRIGGER fail_insert')
|
|
state.reconcile([], [fill])
|
|
self.assertEqual(TState(self.path).items['600000.SH'].base_qty, 100)
|
|
|
|
def test_schema_and_unique_execution(self):
|
|
state = TState(self.path)
|
|
state.reconcile([], [self.deal('base', 'd1', 100, 10)])
|
|
with closing(sqlite3.connect(self.path)) as db:
|
|
tables = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='table'")}
|
|
self.assertEqual(tables, {'positions', 'deals', 'sqlite_sequence'})
|
|
columns = {row[1] for row in db.execute('PRAGMA table_info(deals)')}
|
|
self.assertFalse({'kind', 'confirmed_date', 'deal_ids', 'deal_id', 'order_id', 'qty', 'filled_qty', 'filled_cost', 'amount', 'trade_date', 'trade_time'} & columns)
|
|
self.assertTrue({'sys_order_id', 'local_order_id'} <= columns)
|
|
indexes = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='index'")}
|
|
self.assertTrue({'idx_positions_code', 'idx_positions_base_order_id',
|
|
'idx_positions_added_order_id', 'idx_deals_local_order_id',
|
|
'idx_deals_code_date', 'idx_deals_date_time'} <= indexes)
|
|
self.assertNotIn('kind', state.deals[0])
|
|
self.assertEqual(TState(self.path).deals, state.deals)
|
|
state.deals.append(dict(state.deals[0]))
|
|
with self.assertRaises(sqlite3.IntegrityError):
|
|
state.save()
|
|
self.assertEqual(len(state.deals), 1)
|
|
state.items['600000.SH'].base_cost = float('inf')
|
|
with self.assertRaises(ValueError):
|
|
state.save()
|
|
self.assertEqual(state.items['600000.SH'].base_cost, 10)
|
|
|
|
def test_position_columns_defaults_indexes_and_stable_id(self):
|
|
store = OrderBook(self.path)
|
|
with closing(sqlite3.connect(self.path)) as db:
|
|
columns = [row[1] for row in db.execute('PRAGMA table_info(positions)')]
|
|
self.assertEqual(columns, [
|
|
'id', 'code', 'base_order_id', 'base_qty', 'base_cost',
|
|
'added_order_id', 'added_num', 'added_qty', 'added_cost', 'status',
|
|
])
|
|
indexes = {row[1] for row in db.execute('PRAGMA index_list(positions)')}
|
|
self.assertEqual(indexes, {
|
|
'idx_positions_code', 'idx_positions_base_order_id', 'idx_positions_added_order_id',
|
|
})
|
|
store.save({'600000.SH': {'code': '600000.SH', 'status': READY}}, [])
|
|
position = store.positions['600000.SH']
|
|
first_id = position['id']
|
|
self.assertGreater(first_id, 0)
|
|
self.assertEqual(position['base_order_id'], '')
|
|
self.assertEqual(position['base_qty'], 0)
|
|
self.assertEqual(position['added_cost'], 0.0)
|
|
position.update(base_order_id='base1', base_qty=200, base_cost=10.5,
|
|
added_order_id='add1', added_num=1, added_qty=100, added_cost=9.0,
|
|
status='ACTIVE')
|
|
store.save({'600000.SH': position}, [])
|
|
self.assertEqual(store.positions['600000.SH'], position)
|
|
store.save({}, [])
|
|
store.save({'600001.SH': {'code': '600001.SH', 'status': READY}}, [])
|
|
self.assertGreater(store.positions['600001.SH']['id'], first_id)
|
|
|
|
def test_base_split_fills_and_snapshot_do_not_double_count(self):
|
|
state = TState(self.path)
|
|
first = self.deal('base', 'd1', 40, 10)
|
|
state.reconcile([PositionItem(stock_code='600000.SH', volume=40, open_price=10)], [first])
|
|
second = self.deal('base', 'd2', 60, 12)
|
|
state.reconcile([PositionItem(stock_code='600000.SH', volume=100, open_price=11.2)], [first, second])
|
|
self.assertEqual(state.items['600000.SH'].base_qty, 100)
|
|
self.assertAlmostEqual(state.items['600000.SH'].base_cost, 11.2)
|
|
self.assertEqual(len(state.deals), 2)
|
|
|
|
def test_date_normalization_and_unrelated_strategy(self):
|
|
state = TState(self.path)
|
|
first = self.deal('base', 'd1', 100, 10, '20260901')
|
|
other = self.deal('base', 'd2', 100, 10)
|
|
other.local_order_id = 'trend-base-order'
|
|
state.reconcile([], [first, other])
|
|
first.insert_date = '2026-09-01'
|
|
state.reconcile([], [first])
|
|
self.assertEqual(len(state.deals), 1)
|
|
self.assertEqual(state.deals[0]['insert_date'], '2026-09-01')
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|