184 lines
9.0 KiB
Python
184 lines
9.0 KiB
Python
import sqlite3
|
|
import tempfile
|
|
import unittest
|
|
from contextlib import closing
|
|
from pathlib import Path
|
|
|
|
from libs.state import State
|
|
from sdk import DealItem, PositionItem
|
|
|
|
|
|
class ArchivingTests(unittest.TestCase):
|
|
def setUp(self):
|
|
tmp = tempfile.TemporaryDirectory()
|
|
self.addCleanup(tmp.cleanup)
|
|
self.book = State(Path(tmp.name) / 'state.db')
|
|
|
|
def insert_deal(self, order, qty, amount, time, code='600000.SH', flag=48):
|
|
with closing(self.book._connect()) as db, db:
|
|
db.execute(
|
|
'INSERT INTO deals (stock_code, order_sys_id, order_local_id, offset_flag, '
|
|
'price, volume, trade_amount, trade_date, trade_time) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)',
|
|
(code, order, order, flag, amount / qty, qty, amount, '2026-09-08', time),
|
|
)
|
|
|
|
def test_accumulates_once_and_preserves_base(self):
|
|
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
|
original = dict(self.book.state['600000.SH'])
|
|
self.insert_deal('first', 40, 400, '10:00:00')
|
|
self.insert_deal('second', 60, 720, '10:01:00')
|
|
self.book.archiving()
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual(row['added_qty'], 100)
|
|
self.assertAlmostEqual(row['added_price'], 11.2)
|
|
self.assertEqual(row['added_order_local_id'], 'second')
|
|
self.assertEqual(row['added_created_at'], '2026-09-08 10:01:00')
|
|
for key in original:
|
|
if not key.startswith('added_'):
|
|
self.assertEqual(row[key], original[key])
|
|
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
|
|
restarted = State(self.book.path)
|
|
restarted.archiving()
|
|
self.assertEqual(restarted.state, self.book.state)
|
|
self.insert_deal('late', 50, 500, '09:59:00')
|
|
self.book.archiving()
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual(row['added_qty'], 150)
|
|
self.assertAlmostEqual(row['added_price'], 1620 / 150)
|
|
self.assertEqual(row['added_order_local_id'], 'second')
|
|
|
|
def test_sell_added_then_clear_base(self):
|
|
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
|
self.insert_deal('buy1', 40, 400, '09:59:00')
|
|
self.insert_deal('buy', 60, 600, '10:00:00')
|
|
self.insert_deal('partial', 40, 480, '10:01:00', flag=49)
|
|
self.book.archiving()
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual((row['base_qty'], row['base_price'], row['added_qty'], row['added_price']),
|
|
(100, 8, 60, 10))
|
|
self.insert_deal('sell_added', 60, 720, '10:02:00', flag=49)
|
|
self.book.archiving()
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual((row['base_qty'], row['added_qty'], row['added_price']), (100, 0, 0))
|
|
self.assertEqual((row['added_order_local_id'], row['added_created_at']), ('', ''))
|
|
self.insert_deal('sell_base', 100, 1200, '10:03:00', flag=49)
|
|
self.book.archiving()
|
|
self.assertEqual(self.book.state, {})
|
|
self.book.archiving()
|
|
self.assertEqual(State(self.book.path).state, {})
|
|
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
|
|
|
|
def test_sell_crosses_into_base_then_liquidates(self):
|
|
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
|
self.insert_deal('buy', 50, 500, '10:00:00')
|
|
self.insert_deal('sell', 80, 960, '10:01:00', flag=49)
|
|
self.book.archiving()
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual((row['base_qty'], row['base_price'], row['added_qty']), (70, 8, 0))
|
|
self.insert_deal('buy_again', 30, 300, '10:02:00')
|
|
self.insert_deal('sell_all', 100, 1200, '10:03:00', flag=49)
|
|
self.book.archiving()
|
|
self.assertEqual(self.book.state, {})
|
|
|
|
def test_excess_sell_rolls_back_and_other_directions_are_skipped(self):
|
|
self.insert_deal('other', 100, 1000, '09:59:00', flag=23)
|
|
self.book.archiving()
|
|
self.assertEqual(self.book.state, {})
|
|
self.assertEqual(self.book.deals['other']['is_arch'], 0)
|
|
self.insert_deal('buy', 50, 500, '10:00:00')
|
|
self.insert_deal('sell', 100, 1200, '10:01:00', flag=49)
|
|
errors = self.book.archiving()
|
|
self.assertIn('600000.SH', errors)
|
|
restarted = State(self.book.path)
|
|
self.assertEqual(restarted.state, {})
|
|
self.assertTrue(all(deal['is_arch'] == 0 for deal in restarted.deals.values()))
|
|
|
|
def test_new_state_and_failed_mark_roll_back_together(self):
|
|
self.insert_deal('first', 100, 1000, '10:00:00')
|
|
self.insert_deal('second', 100, 1200, '10:01:00', code='600001.SH')
|
|
self.book.load()
|
|
with closing(self.book._connect()) as db:
|
|
db.execute("""CREATE TRIGGER fail_archive BEFORE UPDATE OF is_arch ON deals
|
|
WHEN OLD.order_sys_id = 'second'
|
|
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
|
|
errors = self.book.archiving()
|
|
self.assertIn('600001.SH', errors)
|
|
restarted = State(self.book.path)
|
|
self.assertEqual(set(restarted.state), {'600000.SH'})
|
|
self.assertEqual(restarted.state, self.book.state)
|
|
self.assertEqual(restarted.deals['first']['is_arch'], 1)
|
|
self.assertEqual(restarted.deals['second']['is_arch'], 0)
|
|
with closing(self.book._connect()) as db:
|
|
db.execute('DROP TRIGGER fail_archive')
|
|
self.book.archiving()
|
|
self.assertEqual(len(self.book.state), 2)
|
|
self.assertEqual(self.book.state['600000.SH']['base_qty'], 0)
|
|
self.assertEqual(self.book.state['600000.SH']['added_qty'], 100)
|
|
|
|
def test_snapshot_matching_buy_is_not_added(self):
|
|
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
|
self.book.sync_deals([DealItem(
|
|
stock_code='600000.SH', order_sys_id='first', remark='base1|test',
|
|
offset_flag=48, volume=100, price=8, trade_amount=800,
|
|
trade_date='20260908', trade_time='100000',
|
|
)])
|
|
self.assertEqual(self.book.archiving(), {})
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual((row['base_qty'], row['added_qty']), (100, 0))
|
|
self.assertEqual(self.book.deals['first']['is_arch'], 1)
|
|
restarted = State(self.book.path)
|
|
self.assertEqual(restarted.archiving(), {})
|
|
self.assertEqual(restarted.state, self.book.state)
|
|
self.insert_deal('new_buy', 100, 1000, '10:01:00')
|
|
with closing(self.book._connect()) as db, db:
|
|
db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'")
|
|
self.assertEqual(self.book.archiving(), {})
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 100, 'CUSTOM'))
|
|
|
|
def test_old_archived_history_without_baseline_is_not_reapplied(self):
|
|
self.insert_deal('old', 100, 1000, '10:00:00')
|
|
with closing(self.book._connect()) as db, db:
|
|
db.execute('UPDATE deals SET is_arch = 1')
|
|
self.insert_deal('new', 50, 500, '10:01:00')
|
|
errors = self.book.archiving()
|
|
self.assertIn('baseline', errors['600000.SH'])
|
|
self.assertEqual(self.book.deals['old']['is_arch'], 1)
|
|
self.assertEqual(self.book.deals['new']['is_arch'], 0)
|
|
|
|
def test_late_buy_replays_after_liquidation_and_restart(self):
|
|
self.insert_deal('buy', 100, 1000, '10:00:00')
|
|
self.insert_deal('sell', 100, 1500, '10:02:00', flag=49)
|
|
self.book.archiving()
|
|
self.assertEqual(self.book.state, {})
|
|
self.book = State(self.book.path)
|
|
self.insert_deal('late', 100, 2000, '100100')
|
|
self.assertEqual(self.book.archiving(), {})
|
|
row = self.book.state['600000.SH']
|
|
self.assertEqual(row['added_qty'], 100)
|
|
self.assertEqual(row['added_price'], 15)
|
|
|
|
def test_snapshot_deletion_before_sell_archiving(self):
|
|
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
|
self.insert_deal('sell', 100, 1000, '10:01:00', flag=49)
|
|
self.book.sync_state([])
|
|
self.book = State(self.book.path)
|
|
self.assertEqual(self.book.archiving(), {})
|
|
self.assertEqual(self.book.state, {})
|
|
self.assertEqual(self.book.deals['sell']['is_arch'], 1)
|
|
|
|
def test_bad_stock_does_not_block_good_stock_and_can_retry(self):
|
|
self.insert_deal('bad_sell', 100, 1500, '10:02:00', flag=49)
|
|
self.insert_deal('good_buy', 100, 1000, '10:00:00', code='600001.SH')
|
|
self.assertIn('600000.SH', self.book.archiving())
|
|
self.assertEqual(self.book.deals['good_buy']['is_arch'], 1)
|
|
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 0)
|
|
self.insert_deal('late_buy', 100, 1000, '10:01:00')
|
|
self.assertEqual(self.book.archiving(), {})
|
|
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 1)
|
|
self.assertNotIn('600000.SH', self.book.state)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|