421 lines
19 KiB
Python
421 lines
19 KiB
Python
"""ZT 轮次状态:幂等成交累计、阶段推进、跨日配额、超期放弃、持久化。"""
|
||
|
||
import json
|
||
import tempfile
|
||
import unittest
|
||
from pathlib import Path
|
||
|
||
from sdk import DealItem, OrderItem, PositionItem
|
||
from strategy.zt.rounds import (
|
||
BASE_SOURCE_OPENED,
|
||
KIND_LONG_T,
|
||
KIND_SHORT_T,
|
||
OUTCOME_ABORTED,
|
||
OUTCOME_BASE,
|
||
OUTCOME_EXPIRED,
|
||
OUTCOME_NORMAL,
|
||
PHASE_CLOSED,
|
||
PHASE_CLOSING,
|
||
PHASE_IDLE,
|
||
PHASE_OPEN,
|
||
PHASE_OPENING,
|
||
Round,
|
||
RoundStore,
|
||
RoundStoreError,
|
||
advance,
|
||
apply_deals,
|
||
expire,
|
||
in_flight_order_ids,
|
||
is_owned_base,
|
||
new_base_round,
|
||
new_round,
|
||
start_round,
|
||
)
|
||
|
||
TODAY = '2026-09-15'
|
||
|
||
|
||
def deal(order_sys_id, remark, volume=100, price=10.0):
|
||
return DealItem(stock_code='600000.SH', order_sys_id=order_sys_id, remark=remark,
|
||
offset_flag=48, volume=volume, price=price,
|
||
trade_amount=price * volume,
|
||
trade_date='20260915', trade_time='100000')
|
||
|
||
|
||
def order(local_id, status):
|
||
return OrderItem(stock_code='600000.SH', order_sys_id=local_id, remark=local_id,
|
||
offset_flag=48, order_status=status,
|
||
insert_date='20260915', insert_time='100000')
|
||
|
||
|
||
class RoundModelTests(unittest.TestCase):
|
||
def test_directions_are_mirrored_between_long_and_short_t(self):
|
||
long_t = Round(code='600000.SH', kind=KIND_LONG_T)
|
||
short_t = Round(code='600000.SH', kind=KIND_SHORT_T)
|
||
self.assertEqual((long_t.entry_side, long_t.exit_side), ('BUY', 'SELL'))
|
||
self.assertEqual((short_t.entry_side, short_t.exit_side), ('SELL', 'BUY'))
|
||
|
||
def test_residual_and_average_prices(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T,
|
||
entry_filled_qty=200, entry_amount=2000.0,
|
||
exit_filled_qty=100, exit_amount=1100.0)
|
||
self.assertEqual(item.residual_qty, 100)
|
||
self.assertAlmostEqual(item.entry_avg_price, 10.0)
|
||
self.assertAlmostEqual(item.exit_avg_price, 11.0)
|
||
self.assertEqual(Round().entry_avg_price, 0.0)
|
||
|
||
def test_daily_quota_and_cross_day_recovery(self):
|
||
item = Round(code='600000.SH')
|
||
self.assertTrue(item.can_open(TODAY))
|
||
item.open_date = TODAY # 今天已开过一轮
|
||
self.assertFalse(item.can_open(TODAY))
|
||
item.open_date = '2026-09-14'
|
||
item.phase = PHASE_OPEN # 昨日未平的轮次继续持有
|
||
self.assertFalse(item.can_open(TODAY))
|
||
item.phase = PHASE_CLOSED
|
||
self.assertTrue(item.can_open(TODAY))
|
||
item.last_trade_date = TODAY # 今天已有腿成交
|
||
self.assertFalse(item.can_open(TODAY))
|
||
item.last_trade_date = '2026-09-14' # 昨日成交,今天可以做一轮
|
||
self.assertTrue(item.can_open(TODAY))
|
||
|
||
def test_new_round_keeps_the_established_base(self):
|
||
item = new_round('600000.SH', KIND_LONG_T, TODAY, 500, 26.89,
|
||
base_date='2026-09-10', base_source='opened')
|
||
self.assertEqual(item.phase, PHASE_OPENING)
|
||
self.assertEqual((item.base_qty, item.base_cost), (500, 26.89))
|
||
self.assertEqual((item.base_date, item.base_source), ('2026-09-10', 'opened'))
|
||
|
||
|
||
class ApplyDealsTests(unittest.TestCase):
|
||
def test_repeated_sync_never_double_counts(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T, entry_order_id='zt-base-1')
|
||
batch = [deal('s1', 'zt-base-1'), deal('s2', 'zt-base-1')]
|
||
apply_deals(item, batch, TODAY)
|
||
self.assertEqual(item.entry_filled_qty, 200)
|
||
self.assertAlmostEqual(item.entry_amount, 2000.0)
|
||
apply_deals(item, batch, TODAY) # 同一批再次同步
|
||
self.assertEqual(item.entry_filled_qty, 200)
|
||
apply_deals(item, batch + [deal('s3', 'zt-base-1')], TODAY)
|
||
self.assertEqual(item.entry_filled_qty, 300)
|
||
self.assertEqual(item.last_trade_date, TODAY)
|
||
|
||
def test_only_this_rounds_legs_are_counted(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T,
|
||
entry_order_id='zt-base-1', exit_order_id='zt-SELL-1')
|
||
apply_deals(item, [
|
||
deal('s1', 'zt-base-1'),
|
||
deal('s2', ''), # 手工单
|
||
deal('s3', 'TREN-BUY-1'), # 其他策略
|
||
deal('s4', 'zt-added-other'), # 本策略但不是本轮
|
||
deal('s5', 'zt-SELL-1', price=11.0),
|
||
], TODAY)
|
||
self.assertEqual(item.entry_filled_qty, 100)
|
||
self.assertEqual(item.exit_filled_qty, 100)
|
||
self.assertEqual(item.residual_qty, 0)
|
||
self.assertEqual(sorted(item.seen_deal_ids), ['s1', 's5'])
|
||
|
||
|
||
class AdvanceTests(unittest.TestCase):
|
||
def test_entry_fully_filled_moves_to_open(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T,
|
||
phase=PHASE_OPENING, entry_order_id='zt-base-1',
|
||
entry_filled_qty=300)
|
||
advance(item, {'other'}, TODAY)
|
||
self.assertEqual(item.phase, PHASE_OPEN)
|
||
self.assertEqual(item.residual_qty, 300)
|
||
|
||
def test_entry_still_in_flight_does_not_move(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T,
|
||
phase=PHASE_OPENING, entry_order_id='zt-base-1',
|
||
entry_filled_qty=100)
|
||
advance(item, {'zt-base-1'}, TODAY)
|
||
self.assertEqual(item.phase, PHASE_OPENING)
|
||
|
||
def test_aborted_entry_releases_the_daily_quota(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T, phase=PHASE_OPENING,
|
||
open_date=TODAY, entry_order_id='zt-base-1')
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.phase, PHASE_CLOSED)
|
||
self.assertEqual(item.outcome, OUTCOME_ABORTED)
|
||
self.assertEqual(item.open_date, '')
|
||
self.assertTrue(item.can_open(TODAY))
|
||
|
||
def test_partially_closed_round_returns_to_open(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T, phase=PHASE_CLOSING,
|
||
entry_order_id='zt-base-1', exit_order_id='zt-SELL-1',
|
||
entry_filled_qty=300, exit_filled_qty=100)
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.phase, PHASE_OPEN)
|
||
self.assertEqual(item.residual_qty, 200)
|
||
|
||
def test_fully_closed_round_finishes(self):
|
||
item = Round(code='600000.SH', kind=KIND_SHORT_T, phase=PHASE_CLOSING,
|
||
open_date=TODAY, entry_order_id='zt-SELL-1', exit_order_id='zt-added-1',
|
||
entry_filled_qty=100, exit_filled_qty=100)
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.phase, PHASE_CLOSED)
|
||
self.assertEqual(item.outcome, OUTCOME_NORMAL)
|
||
self.assertEqual(item.close_date, TODAY)
|
||
self.assertEqual(item.open_date, TODAY) # 完成轮次占用当日配额
|
||
|
||
def test_overnight_round_keeps_its_open_date(self):
|
||
item = Round(code='600000.SH', kind=KIND_SHORT_T, phase=PHASE_OPEN,
|
||
open_date='2026-09-14', entry_order_id='zt-SELL-1',
|
||
entry_filled_qty=100)
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.phase, PHASE_OPEN) # 仍待买回,允许隔夜
|
||
self.assertFalse(item.can_open(TODAY))
|
||
|
||
|
||
class BaseEstablishmentTests(unittest.TestCase):
|
||
def test_base_cost_comes_from_the_actual_fill(self):
|
||
item = new_base_round('600000.SH', TODAY, 300)
|
||
item.entry_order_id = 'zt-base-1'
|
||
apply_deals(item, [deal('s1', 'zt-base-1', volume=300, price=26.89)], TODAY)
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.phase, PHASE_CLOSED)
|
||
self.assertEqual(item.outcome, OUTCOME_BASE)
|
||
self.assertEqual(item.base_qty, 300)
|
||
self.assertAlmostEqual(item.base_cost, 26.89)
|
||
self.assertEqual(item.base_source, BASE_SOURCE_OPENED)
|
||
self.assertEqual(item.base_date, TODAY)
|
||
|
||
def test_partial_base_fill_is_accepted(self):
|
||
item = new_base_round('600000.SH', TODAY, 300)
|
||
item.entry_order_id = 'zt-base-1'
|
||
apply_deals(item, [deal('s1', 'zt-base-1', volume=100, price=26.0)], TODAY)
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual((item.base_qty, item.base_cost), (100, 26.0))
|
||
|
||
def test_empty_base_fill_aborts_and_frees_the_quota(self):
|
||
item = new_base_round('600000.SH', TODAY, 300)
|
||
item.entry_order_id = 'zt-base-1'
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.outcome, OUTCOME_ABORTED)
|
||
self.assertEqual(item.open_date, '')
|
||
self.assertTrue(item.can_open(TODAY))
|
||
|
||
def test_base_round_settles_only_after_the_order_is_no_longer_in_flight(self):
|
||
item = new_base_round('600000.SH', TODAY, 300)
|
||
item.entry_order_id = 'zt-base-1'
|
||
apply_deals(item, [deal('s1', 'zt-base-1', volume=300, price=26.89)], TODAY)
|
||
advance(item, {'zt-base-1'}, TODAY)
|
||
self.assertEqual(item.phase, PHASE_OPENING)
|
||
self.assertEqual(item.base_qty, 0)
|
||
|
||
def test_adoption_is_not_supported(self):
|
||
# 程序不接管账户已有持仓:Round 只认识自己建仓写下的基准。
|
||
item = Round(code='600000.SH', base_qty=500, base_cost=37.72,
|
||
base_source=BASE_SOURCE_OPENED)
|
||
self.assertTrue(is_owned_base(item))
|
||
for source in ('', 'adopted', 'configured'):
|
||
with self.subTest(source=source):
|
||
self.assertFalse(is_owned_base(Round(code='600000.SH', base_qty=500,
|
||
base_cost=37.72,
|
||
base_source=source)))
|
||
self.assertFalse(is_owned_base(Round(code='600000.SH')))
|
||
|
||
def test_apply_deals_reports_applied_fills_for_logging(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T,
|
||
entry_order_id='zt-entry-1', exit_order_id='zt-exit-1')
|
||
applied = apply_deals(item, [deal('s1', 'zt-entry-1'),
|
||
deal('s2', 'TREN-BUY-1'),
|
||
deal('s3', 'zt-exit-1')], TODAY)
|
||
self.assertEqual([leg for leg, _ in applied], ['entry', 'exit'])
|
||
self.assertEqual([entry.order_sys_id for _, entry in applied], ['s1', 's3'])
|
||
self.assertEqual(apply_deals(item, [deal('s1', 'zt-entry-1')], TODAY), [])
|
||
|
||
|
||
class StartRoundTests(unittest.TestCase):
|
||
"""开新轮必须清空上一轮的两条腿,否则残量会静默把本轮判成作废。"""
|
||
|
||
def closed_round(self):
|
||
item = Round(code='600000.SH', kind=KIND_SHORT_T, phase=PHASE_CLOSED,
|
||
open_date='2026-09-14', close_date='2026-09-14',
|
||
outcome=OUTCOME_NORMAL, note='旧备注',
|
||
entry_order_id='e1', entry_plan_qty=500, entry_filled_qty=500,
|
||
entry_amount=5500.0, exit_order_id='x1', exit_plan_qty=500,
|
||
exit_filled_qty=500, exit_amount=4900.0,
|
||
seen_deal_ids=['s1', 's2'])
|
||
return item
|
||
|
||
def test_start_round_clears_both_legs_and_audit_fields(self):
|
||
item = self.closed_round()
|
||
start_round(item, KIND_LONG_T, TODAY)
|
||
self.assertEqual(item.phase, PHASE_OPENING)
|
||
self.assertEqual(item.kind, KIND_LONG_T)
|
||
self.assertEqual(item.open_date, TODAY)
|
||
self.assertEqual((item.close_date, item.outcome, item.note), ('', '', ''))
|
||
self.assertEqual((item.entry_order_id, item.exit_order_id), ('', ''))
|
||
self.assertEqual((item.entry_filled_qty, item.exit_filled_qty), (0, 0))
|
||
self.assertEqual((item.entry_amount, item.exit_amount), (0.0, 0.0))
|
||
self.assertEqual(item.seen_deal_ids, [])
|
||
self.assertEqual(item.residual_qty, 0)
|
||
|
||
def test_start_round_keeps_the_established_base(self):
|
||
item = self.closed_round()
|
||
item.base_qty, item.base_cost = 1000, 10.0
|
||
item.base_date, item.base_source = '2026-09-10', BASE_SOURCE_OPENED
|
||
start_round(item, KIND_SHORT_T, TODAY)
|
||
self.assertEqual((item.base_qty, item.base_cost), (1000, 10.0))
|
||
self.assertEqual((item.base_date, item.base_source),
|
||
('2026-09-10', BASE_SOURCE_OPENED))
|
||
|
||
def test_stale_exit_counter_cannot_abort_a_new_round(self):
|
||
# 复现:直接改字段开新轮,上一轮的 exit_filled_qty 让 residual 变负,
|
||
# advance 会判成作废并立刻重开一轮。
|
||
item = self.closed_round()
|
||
item.kind = KIND_LONG_T
|
||
item.phase = PHASE_OPENING
|
||
item.open_date = TODAY
|
||
item.entry_order_id = 'e2'
|
||
item.entry_filled_qty, item.entry_amount = 0, 0.0
|
||
self.assertEqual(item.residual_qty, -500)
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.phase, PHASE_CLOSED)
|
||
self.assertEqual(item.note, '成交累计异常:平仓量超过开仓量,本轮作废')
|
||
|
||
fixed = self.closed_round()
|
||
start_round(fixed, KIND_LONG_T, TODAY)
|
||
fixed.entry_order_id = 'e2'
|
||
apply_deals(fixed, [deal('s9', 'e2', volume=300, price=9.0)], TODAY)
|
||
advance(fixed, set(), TODAY)
|
||
self.assertEqual(fixed.phase, PHASE_OPEN)
|
||
self.assertEqual(fixed.residual_qty, 300)
|
||
|
||
def test_new_base_round_starts_clean(self):
|
||
item = new_base_round('600000.SH', TODAY, 300)
|
||
self.assertEqual(item.phase, PHASE_OPENING)
|
||
self.assertEqual(item.entry_plan_qty, 300)
|
||
self.assertEqual(item.base_qty, 0)
|
||
|
||
|
||
class ResidualAbsorptionTests(unittest.TestCase):
|
||
"""超期放弃必须把敞口并回底仓,否则会在裸敞口上继续开新轮。"""
|
||
|
||
def test_unclosed_short_t_leg_reduces_the_base(self):
|
||
item = Round(code='600000.SH', kind=KIND_SHORT_T, phase=PHASE_OPEN,
|
||
open_date='2026-09-09', base_qty=1000, base_cost=10.0,
|
||
entry_filled_qty=500, entry_amount=5500.0)
|
||
self.assertTrue(expire(item, TODAY, 5))
|
||
self.assertEqual(item.base_qty, 500) # 卖出未买回,底仓变 500
|
||
self.assertAlmostEqual(item.base_cost, 10.0) # 成本仍是建仓价
|
||
self.assertEqual(item.residual_qty, 500) # 敞口数值保留在审计字段里
|
||
|
||
def test_unclosed_long_t_leg_increases_the_base(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T, phase=PHASE_OPEN,
|
||
open_date='2026-09-09', base_qty=1000, base_cost=10.0,
|
||
entry_filled_qty=300, entry_amount=2700.0)
|
||
self.assertTrue(expire(item, TODAY, 5))
|
||
self.assertEqual(item.base_qty, 1300)
|
||
|
||
def test_normal_completion_leaves_the_base_untouched(self):
|
||
item = Round(code='600000.SH', kind=KIND_SHORT_T, phase=PHASE_CLOSING,
|
||
open_date=TODAY, base_qty=1000, base_cost=10.0,
|
||
entry_order_id='e1', exit_order_id='x1',
|
||
entry_filled_qty=500, entry_amount=5500.0,
|
||
exit_filled_qty=500, exit_amount=4900.0)
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.phase, PHASE_CLOSED)
|
||
self.assertEqual(item.base_qty, 1000)
|
||
|
||
def test_aborted_round_never_touches_the_base(self):
|
||
item = Round(code='600000.SH', kind=KIND_LONG_T, phase=PHASE_OPENING,
|
||
open_date=TODAY, base_qty=1000, base_cost=10.0,
|
||
entry_order_id='e1')
|
||
advance(item, set(), TODAY)
|
||
self.assertEqual(item.outcome, OUTCOME_ABORTED)
|
||
self.assertEqual(item.base_qty, 1000)
|
||
|
||
|
||
class ExpireTests(unittest.TestCase):
|
||
def test_round_beyond_max_hold_days_is_abandoned_not_forced(self):
|
||
item = Round(code='600000.SH', kind=KIND_SHORT_T, phase=PHASE_OPEN,
|
||
open_date='2026-09-09', entry_filled_qty=100)
|
||
self.assertTrue(expire(item, TODAY, 5))
|
||
self.assertEqual(item.phase, PHASE_CLOSED)
|
||
self.assertEqual(item.outcome, OUTCOME_EXPIRED)
|
||
self.assertEqual(item.residual_qty, 100) # 残量留作隔夜,不强平
|
||
|
||
def test_round_within_the_limit_is_kept(self):
|
||
item = Round(code='600000.SH', kind=KIND_SHORT_T, phase=PHASE_OPEN,
|
||
open_date='2026-09-14', entry_filled_qty=100)
|
||
self.assertFalse(expire(item, TODAY, 5))
|
||
self.assertEqual(item.phase, PHASE_OPEN)
|
||
|
||
def test_inactive_rounds_never_expire(self):
|
||
for phase in (PHASE_OPENING, PHASE_CLOSED):
|
||
with self.subTest(phase=phase):
|
||
item = Round(code='600000.SH', phase=phase, open_date='2020-01-01')
|
||
self.assertFalse(expire(item, TODAY, 5))
|
||
|
||
|
||
class InFlightTests(unittest.TestCase):
|
||
def test_only_busy_statuses_count_as_in_flight(self):
|
||
orders = [order(f'zt-o{i}', status) for i, status in
|
||
enumerate(['48', '49', '50', '51', '52', '55', '53', '54', '56', '57'])]
|
||
self.assertEqual(in_flight_order_ids(orders),
|
||
{'zt-o0', 'zt-o1', 'zt-o2', 'zt-o3', 'zt-o4', 'zt-o5'})
|
||
|
||
def test_empty_local_ids_are_ignored(self):
|
||
self.assertEqual(in_flight_order_ids([order('', '50')]), set())
|
||
|
||
|
||
class RoundStoreTests(unittest.TestCase):
|
||
def setUp(self):
|
||
temp = tempfile.TemporaryDirectory()
|
||
self.addCleanup(temp.cleanup)
|
||
self.path = Path(temp.name) / 'zt_rounds.json'
|
||
|
||
def test_roundtrip_survives_restart(self):
|
||
store = RoundStore(self.path)
|
||
item = new_round('600000.SH', KIND_SHORT_T, TODAY, 500, 26.89)
|
||
item.entry_order_id = 'zt-SELL-1'
|
||
item.entry_filled_qty = 300
|
||
item.entry_amount = 8067.0
|
||
item.seen_deal_ids = ['s1', 's2']
|
||
store.put(item)
|
||
store.save()
|
||
|
||
reloaded = RoundStore(self.path)
|
||
restored = reloaded.get('600000.SH')
|
||
self.assertEqual(restored, item)
|
||
self.assertEqual(restored.seen_deal_ids, ['s1', 's2'])
|
||
|
||
def test_missing_file_starts_empty_and_unknown_code_is_idle(self):
|
||
store = RoundStore(self.path)
|
||
self.assertEqual(store.rounds, {})
|
||
self.assertEqual(store.get('600000.SH').phase, PHASE_IDLE)
|
||
|
||
def test_save_leaves_no_temporary_file(self):
|
||
store = RoundStore(self.path)
|
||
store.put(Round(code='600000.SH'))
|
||
store.save()
|
||
self.assertEqual([p.name for p in self.path.parent.iterdir()],
|
||
['zt_rounds.json'])
|
||
|
||
def test_corrupt_or_foreign_state_raises_for_rebuild(self):
|
||
cases = {
|
||
'bad json': '{not json',
|
||
'wrong root': '[]',
|
||
'wrong item': '{"600000.SH": 3}',
|
||
'unknown field': json.dumps({'600000.SH': {'code': '600000.SH', 'zzz': 1}}),
|
||
}
|
||
for label, text in cases.items():
|
||
with self.subTest(label=label):
|
||
self.path.write_text(text, encoding='utf-8')
|
||
with self.assertRaises(RoundStoreError):
|
||
RoundStore(self.path)
|
||
|
||
def test_drop_removes_a_code(self):
|
||
store = RoundStore(self.path)
|
||
store.put(Round(code='600000.SH'))
|
||
store.drop('600000.SH')
|
||
store.drop('600001.SH')
|
||
self.assertEqual(store.rounds, {})
|
||
|
||
|
||
if __name__ == '__main__':
|
||
unittest.main()
|