This commit is contained in:
2026-09-12 16:25:34 +08:00
parent 7a7049ce44
commit 554dd0f4cb
13 changed files with 646 additions and 93 deletions

175
py-client/tests/test_ipo.py Normal file
View File

@@ -0,0 +1,175 @@
import json
import tempfile
import unittest
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from pathlib import Path
from threading import Barrier
from types import SimpleNamespace as NS
from unittest.mock import Mock, patch
from libs.lockfile import claim_json
from sdk import OrderItem
from sdk.trade import TradeMixin
from strategy.ipo import boot
class IPOTests(unittest.TestCase):
def setUp(self):
temp = tempfile.TemporaryDirectory()
self.addCleanup(temp.cleanup)
self.root = Path(temp.name)
self.account = NS(account_id='account-a', enable_auto_ipo=True)
self.global_cfg = NS(qmt_data_dir=temp.name, qmt_base_url='unused', qmt_token='')
self.client = Mock()
self.client.__enter__ = Mock(return_value=self.client)
self.client.__exit__ = Mock(return_value=False)
self.client.orders.return_value = []
self.client.ipo_data.return_value = [dict(stock='600001.SH', issuePrice=10, maxPurchaseNum=100)]
for target, value in [('account_config', self.account), ('global_config', self.global_cfg)]:
ctx = patch.object(boot.config, target, value)
ctx.start()
self.addCleanup(ctx.stop)
ctx = patch.object(boot, 'Client', return_value=self.client)
ctx.start()
self.addCleanup(ctx.stop)
ctx = patch.object(boot, 'datetime')
self.clock = ctx.start()
self.clock.now.return_value = datetime(2026, 9, 11, 10)
self.addCleanup(ctx.stop)
def records(self):
return [json.loads(p.read_text(encoding='utf-8')) for p in sorted(self.root.rglob('*.json'))]
def order(self, status, traded=0):
return OrderItem(stock_code='600001.SH', insert_date='20260911',
remark=self.records()[-1]['order_id'] + '|ipo', offset_flag=23,
order_status=status, volume_traded=traded)
def test_timeout_stays_pending_and_queries_before_next_attempt(self):
self.client.passorder.side_effect = TimeoutError('response lost')
self.assertEqual(boot.AutoBuyIpo(), 0)
self.assertEqual(self.records()[0]['status'], 'pending')
self.assertEqual(boot.AutoBuyIpo(), 0)
self.assertEqual(self.client.passorder.call_count, 1)
self.assertEqual(self.client.orders.call_count, 2)
def test_normal_http_response_is_not_confirmation(self):
self.client.passorder.return_value = {'status': 'success'}
self.assertEqual(boot.AutoBuyIpo(), 1)
self.assertEqual(self.records()[0]['status'], 'pending')
self.assertEqual(boot.AutoBuyIpo(), 0)
self.client.passorder.assert_called_once()
def test_rejection_allows_one_new_attempt_with_new_identity(self):
boot.AutoBuyIpo()
old_id = self.records()[0]['order_id']
self.client.orders.return_value = [self.order(57)]
self.assertEqual(boot.AutoBuyIpo(), 1)
self.assertEqual([r['status'] for r in self.records()], ['rejected', 'pending'])
self.assertNotEqual(self.records()[1]['order_id'], old_id)
boot.AutoBuyIpo() # Old rejection cannot authorize retry of the new attempt.
self.assertEqual(self.client.passorder.call_count, 2)
def test_completed_order_confirms_and_prevents_resubmission(self):
boot.AutoBuyIpo()
self.client.orders.return_value = [self.order(56)]
self.assertEqual(boot.AutoBuyIpo(), 0)
self.assertEqual(self.records()[0]['status'], 'confirmed')
self.client.orders.return_value = []
boot.AutoBuyIpo()
self.client.passorder.assert_called_once()
def test_active_unknown_and_partial_orders_never_retry(self):
boot.AutoBuyIpo()
for status, traded in [(50, 0), (255, 0), (55, 40), (57, 40), (54, 0)]:
with self.subTest(status=status, traded=traded):
self.client.orders.return_value = [self.order(status, traded)]
self.assertEqual(boot.AutoBuyIpo(), 0)
self.client.passorder.assert_called_once()
def test_orders_query_failure_does_not_submit(self):
self.client.orders.side_effect = TimeoutError('unavailable')
self.assertEqual(boot.AutoBuyIpo(), 0)
self.client.passorder.assert_not_called()
self.assertEqual(self.records(), [])
def test_accounts_do_not_share_reservations(self):
boot.AutoBuyIpo()
self.account.account_id = 'account-b'
self.assertEqual(boot.AutoBuyIpo(), 1)
self.assertEqual(self.client.passorder.call_count, 2)
self.assertEqual({r['account'] for r in self.records()}, {'account-a', 'account-b'})
def test_manual_same_day_order_prevents_new_submission(self):
self.client.orders.return_value = [OrderItem(stock_code='600001.SH', offset_flag=23,
order_status=50, insert_date='2026-09-11')]
boot.AutoBuyIpo()
self.client.passorder.assert_not_called()
def test_duplicate_candidates_only_submit_once(self):
self.client.ipo_data.return_value *= 2
self.assertEqual(boot.AutoBuyIpo(), 1)
self.client.passorder.assert_called_once()
def test_concurrent_initial_and_rejected_attempts_are_atomic(self):
for orders in [[], None]:
if orders is None:
orders = [self.order(57)]
barrier = Barrier(2)
def claim(path, record):
barrier.wait(timeout=5)
return claim_json(path, record)
with patch.object(boot, 'claim_json', side_effect=claim), ThreadPoolExecutor(2) as pool:
futures = [pool.submit(boot._subscribe, self.client, orders, 'account-a',
'20260911', '600001.SH', 10, 100) for _ in range(2)]
self.assertEqual(sum(f.result() for f in futures), 1)
self.assertEqual(self.client.passorder.call_count, 2)
def test_pending_record_is_durable_before_request(self):
def submitted(**kwargs):
record = self.records()[0]
self.assertEqual(record['status'], 'pending')
self.assertEqual(record['order_id'], kwargs['order_id'])
self.client.passorder.side_effect = submitted
self.assertEqual(boot.AutoBuyIpo(), 1)
def test_claim_failure_never_submits(self):
with patch.object(boot, 'claim_json', side_effect=OSError('disk failed')):
self.assertEqual(boot.AutoBuyIpo(), 0)
self.client.passorder.assert_not_called()
def test_invalid_candidate_is_skipped_without_blocking_good_one(self):
good = self.client.ipo_data.return_value[0]
invalid = [dict(good, stock='600../x.SH'), dict(good, issuePrice=float('nan')),
dict(good, issuePrice=float('inf')), dict(good, maxPurchaseNum=100.5),
dict(good, maxPurchaseNum=True), dict(good, maxPurchaseNum='NaN'),
dict(good, maxPurchaseNum=0), dict(good, issuePrice=False),
dict(good, stock='600001.SZ'), None]
self.client.ipo_data.return_value = invalid + [good]
self.assertEqual(boot.AutoBuyIpo(), 1)
self.client.passorder.assert_called_once()
self.assertEqual(len(self.records()), 1)
def test_supported_codes_and_numeric_strings(self):
for stock in ['600001.SH', '688001.SH', '689001.SH', '000001.SZ',
'001001.SZ', '002001.SZ', '003001.SZ', '300001.SZ', '301001.SZ']:
self.assertEqual(boot._candidate(dict(stock=stock, issuePrice='10.5', maxPurchaseNum='100')),
(stock, 10.5, 100))
for stock in ['600abc.SH', '600001', '600001.SH/x', '688001.SZ', '300001.SH', None]:
self.assertFalse(boot.is_target_stock(stock))
class IPOResponseTests(unittest.TestCase):
def test_only_list_response_is_accepted(self):
client = TradeMixin()
for value in [None, {}, {'error': 'bad'}, '', 0, False]:
client._post_json = Mock(return_value=value)
with self.subTest(value=value), self.assertRaises(ValueError):
client.ipo_data()
client._post_json = Mock(return_value=[])
self.assertEqual(client.ipo_data(), [])
if __name__ == '__main__':
unittest.main()

View File

@@ -0,0 +1,132 @@
import tempfile
import unittest
from contextlib import closing
from dataclasses import replace
from datetime import datetime, timedelta
from pathlib import Path
from types import SimpleNamespace as NS
from unittest.mock import Mock, patch
from libs.grid_take_profit import GridState
from libs.order import OrderBook
from libs.snapshot import get_collector_snapshot
from libs.state import State
from sdk import Assets, DealItem, OrderItem, PositionItem
from strategy.zt import boot
from strategy.zt.profit import ZTProfitTracker
class ZTAuditFixTests(unittest.TestCase):
def setUp(self):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
self.store = State(Path(tmp.name) / 'state.db')
self.code = '600000.SH'
def deal(self, code, identity):
return DealItem(stock_code=code, order_sys_id=identity,
remark=f'zt-base-{identity}|zt', offset_flag=48,
volume=100, price=10, trade_amount=1000)
def position(self, code):
return PositionItem(stock_code=code, volume=100, open_price=10)
def test_zt_cancels_only_owned_orders_and_tracks_all(self):
stamp = datetime.now() - timedelta(minutes=2)
orders = [OrderItem(stock_code=f'60000{i}.SH', order_sys_id=str(i), remark=remark,
order_status=50, offset_flag=23,
insert_date=stamp.strftime('%Y%m%d'), insert_time=stamp.strftime('%H%M%S'))
for i, remark in enumerate(['zt-base-own|zt', 'zt-SELL-own|zt',
'zt-added-own|zt', 'IPO-new|ipo', '', 'TREN-BUY-other'])]
client = Mock()
book = OrderBook()
book.refresh(client, orders, cancel_prefix='zt-')
self.assertEqual([c.args[0] for c in client.cancel_by_id.call_args_list], ['0', '1', '2'])
self.assertEqual(book.data, orders)
self.assertTrue(all(book.busy(o.stock_code, 'BUY') for o in orders))
client.reset_mock()
book.refresh(client, orders)
self.assertEqual([c.args[0] for c in client.cancel_by_id.call_args_list], ['0', '1', '2', '4', '5'])
def test_invalid_trade_persists_blocks_only_its_stock_and_recovers(self):
good = '600001.SH'
bad = replace(self.deal(self.code, 'bad'), price=float('nan'))
valid = self.deal(good, 'good')
positions = [self.position(c) for c in (self.code, good)]
self.store.sync_account(positions, [bad, valid])
self.assertEqual(self.store.blocked_codes, {self.code})
self.assertEqual(self.store.state[good]['base_qty'], 100)
with closing(self.store._connect()) as db:
payload = db.execute('SELECT payload FROM zt_rejected_deals').fetchone()[0]
self.assertIn('bad', payload)
self.assertIn('NaN', payload)
self.store = State(self.store.path)
self.store.sync_account(positions, [valid])
self.assertEqual(self.store.blocked_codes, {self.code})
self.store.sync_account(positions, [replace(bad, price=10), valid])
self.assertEqual(self.store.blocked_codes, set())
self.assertEqual(self.store.state[self.code]['base_qty'], 100)
def test_invalid_fields_do_not_block_other_stocks(self):
for field, value in [('volume', 0), ('volume', 1.5), ('offset_flag', 99),
('trade_amount', float('inf')), ('price', -1)]:
with self.subTest(field=field):
bad = replace(self.deal(self.code, f'bad-{field}'), **{field: value})
good = self.deal('600001.SH', 'good')
self.store.sync_account([self.position('600001.SH')], [bad, good])
self.assertEqual(self.store.state['600001.SH']['base_qty'], 100)
self.assertIn(self.code, self.store.blocked_codes)
def test_bad_stock_does_not_archive_other_trades_until_corrected(self):
first = self.deal(self.code, 'first')
bad = replace(self.deal(self.code, 'bad'), price=float('nan'))
position = replace(self.position(self.code), volume=200)
self.store.sync_account([position], [first, bad])
self.assertEqual(self.store.deals['first']['is_arch'], 0)
self.assertNotIn(self.code, self.store.state)
self.store.sync_account([position], [replace(bad, price=10)])
self.assertEqual(self.store.state[self.code]['base_qty'], 200)
self.assertEqual(self.store.blocked_codes, set())
def test_profit_basis_changes_reset_peak_but_partial_sell_does_not(self):
tracker = ZTProfitTracker()
position = self.position(self.code)
row = dict(base_qty=100, base_order_local_id='one', base_created_at='now', added_qty=0)
state = NS(blocked_codes=set(), get_by_code=lambda code: row)
tracker.sync_positions([position], state)
self.assertEqual(tracker.observe(self.code, 20).state, GridState.ARMED)
row['base_qty'] = 50
tracker.sync_positions([replace(position, volume=50)], state)
self.assertEqual(tracker.observe(self.code, 19).state, GridState.RETREAT)
for update, cost in [({'added_qty': 100, 'added_price': 11}, 10),
({'added_qty': 0}, 10), ({}, 12),
({'base_order_local_id': 'reopened'}, 12)]:
row.update(update)
tracker.sync_positions([replace(position, open_price=cost)], state)
self.assertEqual(tracker.observe(self.code, 10).state, GridState.ARMED)
tracker.sync_positions([], state)
tracker.sync_positions([position], state)
self.assertEqual(tracker.observe(self.code, 5).state, GridState.ARMED)
def test_run_once_updates_collector_before_market_fetch(self):
positions = [self.position(self.code)]
self.store.sync_account(positions, [], initialize=True)
client = Mock()
run = NS(client=client, account_cfg=NS(account_id='zt-test', min_cash_ratio=0.1),
orders=Mock(), profit_tracker=ZTProfitTracker())
client.deals.return_value = []
client.portfolio.return_value = NS(assets=Assets(20000, 10000),
positions={self.code: positions[0]}, orders=[])
client.full_tick.side_effect = RuntimeError('no market data')
with patch.object(boot, 'trading_time', return_value=True), \
patch.object(boot, 'market_allow_open', return_value=True):
boot.RunOnce(run, self.store, [])
snapshot = get_collector_snapshot()
self.assertEqual(snapshot[0], 'zt-test')
self.assertEqual(snapshot[1].total, 20000)
self.assertEqual(snapshot[2], positions)
run.orders.refresh.assert_called_once_with(client, [], cancel_prefix='zt-')
if __name__ == '__main__':
unittest.main()