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 IPOResponseClient(TradeMixin): """真实解析 + 模拟传输:用于验证 QMT 原始响应到下单的完整链路。""" def __init__(self, payload, orders=None): self.payload = payload self._orders = orders or [] self.submitted = [] def _post_json(self, path, body=None): return self.payload def orders(self): return self._orders def passorder(self, **kwargs): self.submitted.append(kwargs) return {'status': 'success'} def __enter__(self): return self def __exit__(self, *_args): return False 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_code_keyed_qmt_response_submits_subscription(self): client = IPOResponseClient( {'301001.SZ': {'issuePrice': 12.5, 'maxPurchaseNum': 15000}}) with patch.object(boot, 'Client', return_value=client): self.assertEqual(boot.AutoBuyIpo(), 1) self.assertEqual(len(client.submitted), 1) self.assertEqual(client.submitted[0]['stock'], '301001.SZ') self.assertEqual(client.submitted[0]['volume'], 15000) self.assertEqual(client.submitted[0]['price'], 12.5) self.assertEqual([r['status'] for r in self.records()], ['pending']) def test_beijing_candidate_is_excluded_without_error(self): client = IPOResponseClient({ '920202.BJ': {'issuePrice': 7.55, 'maxPurchaseNum': 1190000}, '301716.SZ': {'issuePrice': 10, 'maxPurchaseNum': 3500}, }) with patch.object(boot, 'Client', return_value=client), self.assertLogs(level='INFO') as logs: self.assertEqual(boot.AutoBuyIpo(), 1) self.assertTrue(all(record.levelno == 20 for record in logs.records)) self.assertEqual([order['stock'] for order in client.submitted], ['301716.SZ']) self.assertEqual([record['stock'] for record in self.records()], ['301716.SZ']) 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 parsed(self, value): client = TradeMixin() client._post_json = Mock(return_value=value) return client.ipo_data() def test_code_keyed_mapping_is_parsed(self): payload = { '301001.SZ': {'issuePrice': 12.5, 'maxPurchaseNum': 15000, 'stockName': '示例'}, '601127.SH': {'issuePrice': 4.09, 'maxPurchaseNum': 15000}, } self.assertEqual(self.parsed(payload), [ dict(stock='301001.SZ', issuePrice=12.5, maxPurchaseNum=15000, stockName='示例'), dict(stock='601127.SH', issuePrice=4.09, maxPurchaseNum=15000), ]) def test_bare_code_uses_market_field(self): payload = {'301001': {'market': 'SZ', 'issuePrice': 12.5, 'maxPurchaseNum': 15000}} self.assertEqual(self.parsed(payload), [ dict(stock='301001.SZ', market='SZ', issuePrice=12.5, maxPurchaseNum=15000), ]) def test_market_bucketed_mapping_is_flattened(self): payload = { 'SH': {'601127': {'issuePrice': 4.09, 'maxPurchaseNum': 15000}}, 'SZ': {'301001': {'issuePrice': 12.5, 'maxPurchaseNum': 15000}}, } self.assertEqual(self.parsed(payload), [ dict(stock='601127.SH', market='SH', issuePrice=4.09, maxPurchaseNum=15000), dict(stock='301001.SZ', market='SZ', issuePrice=12.5, maxPurchaseNum=15000), ]) def test_legacy_data_wrapper_is_unwrapped(self): payload = {'data': {'301001.SZ': {'issuePrice': 12.5, 'maxPurchaseNum': 15000}}} self.assertEqual(self.parsed(payload), [ dict(stock='301001.SZ', issuePrice=12.5, maxPurchaseNum=15000), ]) def test_list_response_passes_through(self): payload = [dict(stock='600001.SH', issuePrice=10, maxPurchaseNum=100)] self.assertEqual(self.parsed(payload), payload) def test_empty_response_means_no_candidate(self): for value in [None, {}, [], {'SH': {}, 'SZ': {}}]: with self.subTest(value=value): self.assertEqual(self.parsed(value), []) def test_unrecognised_structure_raises(self): for value in ['', 0, False, 'unexpected', {'error': 'bad'}, {'SH': 'x'}]: with self.subTest(value=value), self.assertRaises(ValueError): self.parsed(value) if __name__ == '__main__': unittest.main()