Files
big-qmt/py-client/tests/test_orderbook.py

216 lines
11 KiB
Python
Raw Normal View History

2026-09-07 00:27:33 +08:00
import sqlite3
import tempfile
import unittest
from contextlib import closing
2026-09-07 14:04:26 +08:00
from dataclasses import asdict, fields
2026-09-07 00:27:33 +08:00
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]
2026-09-07 14:04:26 +08:00
return DealItem(
order_sys_id=sys_order_id, stock_code='600000.SH',
offset_flag=24 if kind == 'sell' else 23,
volume=qty, price=price, trade_amount=qty * price,
trade_date=date, trade_time='10:00:00', remark=prefix + 'order1|zt',
)
2026-09-07 00:27:33 +08:00
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)
2026-09-07 18:18:00 +08:00
book.sync_deals([first, first, second])
2026-09-07 00:27:33 +08:00
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
2026-09-07 14:04:26 +08:00
self.assertEqual(book.deals['d1']['trade_date'], '2026-09-01')
self.assertEqual(book.deals['d2']['volume'], 60)
self.assertEqual(book.deals['d2']['order_local_id'], 'zt-base-order1')
2026-09-07 00:27:33 +08:00
book = OrderBook(self.path)
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
2026-09-07 14:04:26 +08:00
self.assertEqual(book.deals['d1']['order_local_id'], 'zt-base-order1')
2026-09-07 18:18:00 +08:00
with patch.object(book, '_connect') as connect:
2026-09-07 00:27:33 +08:00
book.sync_deals([first, second])
book.sync_deals([])
2026-09-07 18:18:00 +08:00
connect.assert_not_called()
2026-09-07 00:27:33 +08:00
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)
2026-09-07 14:04:26 +08:00
invalid.offset_flag = -1
2026-09-07 00:27:33 +08:00
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, {})
2026-09-07 14:04:26 +08:00
invalid.offset_flag = 23
2026-09-07 00:27:33 +08:00
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)
2026-09-07 14:04:26 +08:00
writer.sync_positions([PositionItem(stock_code='600000.SH', volume=100)])
2026-09-07 00:27:33 +08:00
writer.sync_deals([self.deal('base', 'd1', 100, 10)])
book.load()
2026-09-07 14:04:26 +08:00
self.assertEqual(book.positions['600000.SH']['volume'], 100)
2026-09-07 00:27:33 +08:00
self.assertEqual(book.deals_sys_ids, {'d1'})
2026-09-07 14:04:26 +08:00
self.assertEqual(book.deals['d1']['remark'], 'zt-base-order1|zt')
2026-09-07 00:27:33 +08:00
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)')}
2026-09-07 14:04:26 +08:00
self.assertEqual(columns, {'id', 'order_local_id', *(field.name for field in fields(DealItem))})
2026-09-07 00:27:33 +08:00
indexes = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='index'")}
2026-09-07 14:04:26 +08:00
self.assertTrue({'idx_positions_stock_code', 'idx_deals_order_sys_id',
'idx_deals_order_ref', 'idx_deals_stock_code_date', 'idx_deals_date_time'} <= indexes)
for index, expected in (
('idx_deals_order_sys_id', ['order_sys_id']),
('idx_deals_order_ref', ['order_local_id']),
('idx_deals_stock_code_date', ['stock_code']),
('idx_deals_date_time', ['trade_date']),
):
self.assertEqual([row[2] for row in db.execute(f'PRAGMA index_info({index})')], expected)
2026-09-07 00:27:33 +08:00
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)
2026-09-07 18:18:00 +08:00
store.sync_deals([self.deal('base', 'd1', 100, 10)])
saved_deals = dict(store.deals)
2026-09-07 00:27:33 +08:00
with closing(sqlite3.connect(self.path)) as db:
2026-09-07 14:04:26 +08:00
columns = {row[1] for row in db.execute('PRAGMA table_info(positions)')}
self.assertEqual(columns, {'id', *(field.name for field in fields(PositionItem))})
2026-09-07 00:27:33 +08:00
indexes = {row[1] for row in db.execute('PRAGMA index_list(positions)')}
2026-09-07 14:04:26 +08:00
self.assertEqual(indexes, {'idx_positions_stock_code'})
position = PositionItem(stock_code='600000.SH', volume=100, open_price=10,
stock_name='stock', can_use_volume=100, float_profit=-2.5)
store.sync_positions([position])
saved = store.positions[position.stock_code]
first_id = saved['id']
self.assertEqual({k: v for k, v in saved.items() if k != 'id'}, asdict(position))
position.volume = 200
store.sync_positions([position])
self.assertEqual(store.positions[position.stock_code]['id'], first_id)
self.assertEqual(store.positions[position.stock_code]['volume'], 200)
store.sync_positions([])
self.assertEqual(store.positions, {})
2026-09-07 18:18:00 +08:00
self.assertEqual(store.deals, saved_deals)
2026-09-07 14:04:26 +08:00
store.sync_positions([PositionItem(stock_code='600001.SH')])
2026-09-07 00:27:33 +08:00
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)
2026-09-07 14:04:26 +08:00
other.remark = 'trend-base-order'
2026-09-07 00:27:33 +08:00
state.reconcile([], [first, other])
2026-09-07 14:04:26 +08:00
first.trade_date = '2026-09-01'
2026-09-07 00:27:33 +08:00
state.reconcile([], [first])
self.assertEqual(len(state.deals), 1)
2026-09-07 14:04:26 +08:00
self.assertEqual(state.deals[0]['trade_date'], '2026-09-01')
2026-09-07 00:27:33 +08:00
if __name__ == '__main__':
unittest.main()