Files
big-qmt/py-client/tests/test_orderbook.py
2026-09-07 18:18:00 +08:00

216 lines
11 KiB
Python

import sqlite3
import tempfile
import unittest
from contextlib import closing
from dataclasses import asdict, fields
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(
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',
)
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)
book.sync_deals([first, first, second])
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
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')
book = OrderBook(self.path)
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
self.assertEqual(book.deals['d1']['order_local_id'], 'zt-base-order1')
with patch.object(book, '_connect') as connect:
book.sync_deals([first, second])
book.sync_deals([])
connect.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.offset_flag = -1
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.offset_flag = 23
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.sync_positions([PositionItem(stock_code='600000.SH', volume=100)])
writer.sync_deals([self.deal('base', 'd1', 100, 10)])
book.load()
self.assertEqual(book.positions['600000.SH']['volume'], 100)
self.assertEqual(book.deals_sys_ids, {'d1'})
self.assertEqual(book.deals['d1']['remark'], 'zt-base-order1|zt')
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.assertEqual(columns, {'id', 'order_local_id', *(field.name for field in fields(DealItem))})
indexes = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='index'")}
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)
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)
store.sync_deals([self.deal('base', 'd1', 100, 10)])
saved_deals = dict(store.deals)
with closing(sqlite3.connect(self.path)) as db:
columns = {row[1] for row in db.execute('PRAGMA table_info(positions)')}
self.assertEqual(columns, {'id', *(field.name for field in fields(PositionItem))})
indexes = {row[1] for row in db.execute('PRAGMA index_list(positions)')}
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, {})
self.assertEqual(store.deals, saved_deals)
store.sync_positions([PositionItem(stock_code='600001.SH')])
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.remark = 'trend-base-order'
state.reconcile([], [first, other])
first.trade_date = '2026-09-01'
state.reconcile([], [first])
self.assertEqual(len(state.deals), 1)
self.assertEqual(state.deals[0]['trade_date'], '2026-09-01')
if __name__ == '__main__':
unittest.main()