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

152 lines
6.7 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
2026-09-08 15:18:09 +08:00
from libs.state import State, StateItem
2026-09-07 00:27:33 +08:00
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]
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_json_is_never_read(self):
legacy = self.path.with_suffix('.json')
legacy.write_text('invalid JSON', encoding='utf-8')
2026-09-08 15:18:09 +08:00
book = State(self.path)
2026-09-07 00:27:33 +08:00
self.assertIsNone(book.load())
2026-09-10 12:50:40 +08:00
self.assertEqual((book.state, book.deals, book.deals_sys_ids), ({}, {}, set()))
2026-09-07 00:27:33 +08:00
self.assertEqual(legacy.read_text(encoding='utf-8'), 'invalid JSON')
def test_sync_deals_deduplicates_batch_and_restart(self):
2026-09-08 15:18:09 +08:00
book = State(self.path)
2026-09-10 12:50:40 +08:00
self.assertEqual((book.state, book.deals, book.deals_sys_ids), ({}, {}, set()))
2026-09-07 00:27:33 +08:00
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-08 15:18:09 +08:00
book = State(self.path)
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']['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):
2026-09-08 15:18:09 +08:00
book = State(self.path)
2026-09-07 00:27:33 +08:00
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())
2026-09-08 15:18:09 +08:00
self.assertEqual(State(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):
2026-09-08 15:18:09 +08:00
book = State(self.path)
writer = State(self.path)
writer.sync_state([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-10 12:50:40 +08:00
self.assertEqual(book.state['600000.SH']['base_qty'], 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_position_columns_defaults_indexes_and_stable_id(self):
2026-09-08 15:18:09 +08:00
store = State(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-08 15:18:09 +08:00
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'})
2026-09-07 14:04:26 +08:00
position = PositionItem(stock_code='600000.SH', volume=100, open_price=10,
stock_name='stock', can_use_volume=100, float_profit=-2.5)
2026-09-08 15:18:09 +08:00
store.sync_state([position])
2026-09-10 12:50:40 +08:00
saved = store.state[position.stock_code]
2026-09-07 14:04:26 +08:00
first_id = saved['id']
2026-09-08 15:18:09 +08:00
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'])
2026-09-07 14:04:26 +08:00
position.volume = 200
2026-09-08 15:18:09 +08:00
position.open_price = 12
store.sync_state([position])
2026-09-10 12:50:40 +08:00
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)
2026-09-08 15:18:09 +08:00
store.sync_state([position, PositionItem(stock_code='600001.SH', volume=100)])
2026-09-10 12:50:40 +08:00
self.assertEqual(store.state[position.stock_code], saved)
self.assertEqual(store.state['600001.SH']['base_qty'], 100)
2026-09-08 15:18:09 +08:00
position.volume = 0
store.sync_state([position, PositionItem(stock_code='600002.SH')])
2026-09-10 12:50:40 +08:00
self.assertEqual(store.state, {})
self.assertEqual(State(self.path).state, {})
2026-09-08 15:18:09 +08:00
store.sync_state([PositionItem(stock_code='600001.SH', volume=100)])
2026-09-10 12:50:40 +08:00
self.assertGreater(store.state['600001.SH']['id'], first_id)
2026-09-08 15:18:09 +08:00
store.sync_state([])
2026-09-10 12:50:40 +08:00
self.assertEqual(store.state, {})
2026-09-07 18:18:00 +08:00
self.assertEqual(store.deals, saved_deals)
2026-09-08 15:18:09 +08:00
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',
))
2026-09-10 12:50:40 +08: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']]
2026-09-08 15:18:09 +08:00
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)])
2026-09-10 12:50:40 +08:00
self.assertEqual(book.state[row['stock_code']], saved)
2026-09-08 15:18:09 +08:00
with self.assertRaises(sqlite3.IntegrityError):
2026-09-10 12:50:40 +08:00
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)
2026-09-07 00:27:33 +08:00
if __name__ == '__main__':
unittest.main()