Files
big-qmt/py-client/tests/test_orderbook.py
2026-09-10 12:50:40 +08:00

152 lines
6.7 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.state import State, StateItem
from sdk import DealItem, PositionItem
class OrderBookTests(unittest.TestCase):
def setUp(self):
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_json_is_never_read(self):
legacy = self.path.with_suffix('.json')
legacy.write_text('invalid JSON', encoding='utf-8')
book = State(self.path)
self.assertIsNone(book.load())
self.assertEqual((book.state, 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 = State(self.path)
self.assertEqual((book.state, 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 = State(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 = State(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(State(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 = State(self.path)
writer = State(self.path)
writer.sync_state([PositionItem(stock_code='600000.SH', volume=100)])
writer.sync_deals([self.deal('base', 'd1', 100, 10)])
book.load()
self.assertEqual(book.state['600000.SH']['base_qty'], 100)
self.assertEqual(book.deals_sys_ids, {'d1'})
self.assertEqual(book.deals['d1']['remark'], 'zt-base-order1|zt')
def test_position_columns_defaults_indexes_and_stable_id(self):
store = State(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(state)')}
self.assertEqual(columns, {'id', *(field.name for field in fields(StateItem))})
indexes = {row[1] for row in db.execute('PRAGMA index_list(state)')}
self.assertEqual(indexes, {'idx_state_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_state([position])
saved = store.state[position.stock_code]
first_id = saved['id']
self.assertEqual(saved['base_qty'], 100)
self.assertEqual(saved['base_price'], 10)
self.assertEqual(saved['added_qty'], 0)
self.assertEqual(saved['base_order_local_id'], '')
self.assertTrue(saved['base_created_at'])
position.volume = 200
position.open_price = 12
store.sync_state([position])
self.assertEqual(store.state[position.stock_code]['id'], first_id)
self.assertEqual(store.state[position.stock_code], saved)
self.assertEqual(State(self.path).state[position.stock_code], saved)
store.sync_state([position, PositionItem(stock_code='600001.SH', volume=100)])
self.assertEqual(store.state[position.stock_code], saved)
self.assertEqual(store.state['600001.SH']['base_qty'], 100)
position.volume = 0
store.sync_state([position, PositionItem(stock_code='600002.SH')])
self.assertEqual(store.state, {})
self.assertEqual(State(self.path).state, {})
store.sync_state([PositionItem(stock_code='600001.SH', volume=100)])
self.assertGreater(store.state['600001.SH']['id'], first_id)
store.sync_state([])
self.assertEqual(store.state, {})
self.assertEqual(store.deals, saved_deals)
def test_state_fields_survive_restart_and_sync(self):
book = State(self.path)
row = asdict(StateItem(
stock_code='600000.SH', status='READY',
base_order_local_id='base-1', base_qty=100, base_price=10,
base_created_at='2026-09-08T09:30:00',
added_order_local_id='added-1', added_qty=50, added_price=9,
added_created_at='2026-09-08T10:30:00',
))
with closing(book._connect()) as db, db:
db.execute(
f"INSERT INTO state ({', '.join(row)}) VALUES ({', '.join(':' + key for key in row)})",
row,
)
book.load()
saved = book.state[row['stock_code']]
self.assertEqual({k: v for k, v in saved.items() if k != 'id'}, row)
book = State(self.path)
book.sync_state([PositionItem(stock_code=row['stock_code'], volume=150, open_price=9.5)])
self.assertEqual(book.state[row['stock_code']], saved)
with self.assertRaises(sqlite3.IntegrityError):
with closing(book._connect()) as db, db:
db.execute('UPDATE state SET added_qty = -1')
self.assertEqual(State(self.path).state[row['stock_code']], saved)
if __name__ == '__main__':
unittest.main()