This commit is contained in:
2026-09-15 20:02:05 +08:00
parent bfe89ba122
commit 04daeff141
37 changed files with 2674 additions and 1426 deletions

View File

@@ -0,0 +1,152 @@
"""ZT 正T/反T 规则:方向选择、手数与资金/库存封顶、买卖触发条件。"""
import unittest
from strategy.zt.rounds import KIND_LONG_T, KIND_SHORT_T
from strategy.zt.rules import (
choose_kind,
entry_triggered,
entry_volume,
exit_triggered,
exit_volume,
price_allowed,
)
BASE_COST = 10.0
class ChooseKindTests(unittest.TestCase):
def test_band_decides_the_direction(self):
self.assertEqual(choose_kind(9.0, BASE_COST, 1.0), KIND_LONG_T)
self.assertEqual(choose_kind(11.0, BASE_COST, 1.0), KIND_SHORT_T)
def test_neutral_band_does_nothing(self):
for price in (9.91, 10.0, 10.09):
with self.subTest(price=price):
self.assertIsNone(choose_kind(price, BASE_COST, 1.0))
def test_invalid_inputs_yield_no_direction(self):
for price, cost, band in ((0, BASE_COST, 1.0), (-1, BASE_COST, 1.0),
(9.0, 0, 1.0), (9.0, BASE_COST, -1)):
with self.subTest(price=price, cost=cost, band=band):
self.assertIsNone(choose_kind(price, cost, band))
def test_zero_band_picks_a_side_but_never_both(self):
self.assertEqual(choose_kind(9.99, BASE_COST, 0), KIND_LONG_T)
self.assertEqual(choose_kind(10.01, BASE_COST, 0), KIND_SHORT_T)
def test_price_cap(self):
self.assertTrue(price_allowed(199.0, 200.0))
self.assertFalse(price_allowed(200.01, 200.0))
self.assertFalse(price_allowed(0, 200.0))
class EntryVolumeTests(unittest.TestCase):
def test_long_t_uses_hands_and_is_capped_by_cash(self):
self.assertEqual(entry_volume(KIND_LONG_T, price=10.0, open_hands=3,
sell_ratio=0.5, base_qty=0,
can_use_volume=0, available=100000.0), 300)
self.assertEqual(entry_volume(KIND_LONG_T, price=10.0, open_hands=3,
sell_ratio=0.5, base_qty=0,
can_use_volume=0, available=2500.0), 200)
self.assertEqual(entry_volume(KIND_LONG_T, price=10.0, open_hands=3,
sell_ratio=0.5, base_qty=0,
can_use_volume=0, available=999.0), 0)
def test_long_t_never_forces_a_lot_when_cash_is_short(self):
# 与 calc_buy_volume 的 max(1, ...) 不同:这里买不起就不买。
self.assertEqual(entry_volume(KIND_LONG_T, price=1500.0, open_hands=1,
sell_ratio=0.5, base_qty=0,
can_use_volume=0, available=5000.0), 0)
def test_short_t_uses_ratio_and_is_capped_by_sellable_inventory(self):
self.assertEqual(entry_volume(KIND_SHORT_T, price=10.0, open_hands=3,
sell_ratio=0.5, base_qty=1000,
can_use_volume=1000, available=0.0), 500)
self.assertEqual(entry_volume(KIND_SHORT_T, price=10.0, open_hands=3,
sell_ratio=0.5, base_qty=1000,
can_use_volume=250, available=0.0), 200)
self.assertEqual(entry_volume(KIND_SHORT_T, price=10.0, open_hands=3,
sell_ratio=0.5, base_qty=1000,
can_use_volume=99, available=0.0), 0)
self.assertEqual(entry_volume(KIND_SHORT_T, price=10.0, open_hands=3,
sell_ratio=0.5, base_qty=100,
can_use_volume=100, available=0.0), 0)
def test_unknown_kind_or_bad_price_does_nothing(self):
self.assertEqual(entry_volume('???', price=10.0, open_hands=3, sell_ratio=0.5,
base_qty=100, can_use_volume=100, available=1e6), 0)
self.assertEqual(entry_volume(KIND_LONG_T, price=0.0, open_hands=3, sell_ratio=0.5,
base_qty=0, can_use_volume=0, available=1e6), 0)
class ExitVolumeTests(unittest.TestCase):
def test_long_t_exit_is_limited_by_sellable_inventory(self):
# 正T 当天买入的份额 T+1 才可卖:可卖为 0 时只能留成隔夜。
self.assertEqual(exit_volume(KIND_LONG_T, residual_qty=300, price=10.0,
can_use_volume=0, available=1e6), 0)
self.assertEqual(exit_volume(KIND_LONG_T, residual_qty=300, price=10.0,
can_use_volume=300, available=1e6), 300)
self.assertEqual(exit_volume(KIND_LONG_T, residual_qty=300, price=10.0,
can_use_volume=250, available=1e6), 200)
self.assertEqual(exit_volume(KIND_LONG_T, residual_qty=150, price=10.0,
can_use_volume=100, available=1e6), 100)
def test_short_t_exit_is_limited_by_cash(self):
self.assertEqual(exit_volume(KIND_SHORT_T, residual_qty=500, price=10.0,
can_use_volume=0, available=3000.0), 300)
self.assertEqual(exit_volume(KIND_SHORT_T, residual_qty=500, price=10.0,
can_use_volume=0, available=100000.0), 500)
self.assertEqual(exit_volume(KIND_SHORT_T, residual_qty=500, price=10.0,
can_use_volume=0, available=50.0), 0)
def test_nothing_to_close(self):
for residual in (0, -100):
with self.subTest(residual=residual):
self.assertEqual(exit_volume(KIND_LONG_T, residual_qty=residual, price=10.0,
can_use_volume=1000, available=1e6), 0)
class TriggerTests(unittest.TestCase):
def test_entry_needs_direction_plus_confirmation(self):
self.assertTrue(entry_triggered(KIND_LONG_T, 9.0, BASE_COST, band_pct=1.0,
rebound_confirmed=True, retrace_confirmed=False))
self.assertFalse(entry_triggered(KIND_LONG_T, 9.0, BASE_COST, band_pct=1.0,
rebound_confirmed=False, retrace_confirmed=True))
self.assertTrue(entry_triggered(KIND_SHORT_T, 11.0, BASE_COST, band_pct=1.0,
rebound_confirmed=False, retrace_confirmed=True))
# 方向与位置不符时即使确认也不触发
self.assertFalse(entry_triggered(KIND_SHORT_T, 9.0, BASE_COST, band_pct=1.0,
rebound_confirmed=True, retrace_confirmed=True))
self.assertFalse(entry_triggered(KIND_LONG_T, 10.0, BASE_COST, band_pct=1.0,
rebound_confirmed=True, retrace_confirmed=True))
def test_short_t_exit_needs_fall_and_rebound(self):
kwargs = dict(buy_fall_pct=1.0, profit_step_pct=1.0)
self.assertTrue(exit_triggered(KIND_SHORT_T, 9.8, 10.0,
rebound_confirmed=True, **kwargs))
self.assertFalse(exit_triggered(KIND_SHORT_T, 9.8, 10.0,
rebound_confirmed=False, **kwargs))
self.assertFalse(exit_triggered(KIND_SHORT_T, 9.95, 10.0,
rebound_confirmed=True, **kwargs))
def test_long_t_exit_needs_a_profit_step(self):
kwargs = dict(buy_fall_pct=1.0, profit_step_pct=1.0)
self.assertTrue(exit_triggered(KIND_LONG_T, 10.1, 10.0,
rebound_confirmed=False, **kwargs))
self.assertFalse(exit_triggered(KIND_LONG_T, 10.0, 10.0,
rebound_confirmed=True, **kwargs))
# 正T 平仓不看回落,回落到成本之下不卖
self.assertFalse(exit_triggered(KIND_LONG_T, 9.8, 10.0,
rebound_confirmed=True, **kwargs))
def test_missing_basis_never_triggers(self):
self.assertFalse(exit_triggered(KIND_LONG_T, 12.0, 0.0,
buy_fall_pct=1.0, profit_step_pct=1.0,
rebound_confirmed=True))
self.assertFalse(exit_triggered('???', 12.0, 10.0, buy_fall_pct=1.0,
profit_step_pct=1.0, rebound_confirmed=True))
if __name__ == '__main__':
unittest.main()