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,104 @@
"""ZT 配置:手数、中性带、最长持有天数与开关的校验。"""
import tempfile
import unittest
from pathlib import Path
from unittest.mock import Mock, patch
import yaml
import config
from config import AccountConfig, GlobalConfig, SignalConfig
from strategy.zt import boot
class ZTConfigTests(unittest.TestCase):
def zt_config(self, directory, **overrides):
root = Path(directory)
(root / '_global.yaml').write_text(yaml.safe_dump({
'qmt_base_url': 'unused', 'api_host': 'unused',
'qmt_data_dir': directory, 'hosts': {'test': 'account'},
}), encoding='utf-8')
account = dict(buy_value=1000, strategy='zt', signal_allow=['dcm'])
account.update(overrides)
(root / 'account.yaml').write_text(yaml.safe_dump(account), encoding='utf-8')
return root
def test_config_accepts_only_nonnegative_integer_hands(self):
with tempfile.TemporaryDirectory() as directory, \
patch.object(config, 'global_config'), patch.object(config, 'account_config'):
for hands in (None, 0, 3, -1, 1.5, '3', True):
with self.subTest(hands=hands):
account = dict(buy_value=1000, strategy='zt', signal_allow=['dcm'])
if hands is not None:
account['zt_open_hands'] = hands
root = self.zt_config(directory, **{k: v for k, v in
account.items()
if k not in ('buy_value',
'strategy',
'signal_allow')})
if hands is None or type(hands) is int and hands >= 0:
_, loaded = config.load(root, 'test')
self.assertEqual(loaded.zt_open_hands, hands or 0)
else:
with self.assertRaisesRegex(ValueError, 'zt_open_hands'):
config.load(root, 'test')
def test_t_band_and_hold_days_defaults_and_validation(self):
with tempfile.TemporaryDirectory() as directory, \
patch.object(config, 'global_config'), patch.object(config, 'account_config'):
_, loaded = config.load(self.zt_config(directory), 'test')
self.assertEqual(loaded.zt_t_band_pct, 1.0)
self.assertEqual(loaded.zt_max_hold_days, 5)
for band, valid in ((0, True), (0.5, True), (1.0, True), (-1, False)):
with self.subTest(band=band):
root = self.zt_config(directory, zt_t_band_pct=band)
if valid:
_, loaded = config.load(root, 'test')
self.assertEqual(loaded.zt_t_band_pct, band)
else:
with self.assertRaisesRegex(ValueError, 'zt_t_band_pct'):
config.load(root, 'test')
for days, valid in ((1, True), (5, True), (0, False), (-1, False),
(1.5, False), ('5', False), (True, False)):
with self.subTest(days=days):
root = self.zt_config(directory, zt_max_hold_days=days)
if valid:
_, loaded = config.load(root, 'test')
self.assertEqual(loaded.zt_max_hold_days, days)
else:
with self.assertRaisesRegex(ValueError, 'zt_max_hold_days'):
config.load(root, 'test')
def test_zero_hands_does_not_initialize_strategy(self):
with patch.object(config, 'account_config', AccountConfig()), \
patch.object(boot, 'Client') as client, \
patch.object(boot, '_open_store') as store, \
patch.object(boot, 'init_signals') as signals:
boot.StartZT()
for dependency in (client, store, signals):
dependency.assert_not_called()
def test_unknown_account_key_is_rejected_with_a_clear_error(self):
with tempfile.TemporaryDirectory() as directory, \
patch.object(config, 'global_config'), patch.object(config, 'account_config'):
root = self.zt_config(directory, zt_sell_ratios=0.5) # 拼错的名字
with self.assertRaisesRegex(ValueError, 'zt_sell_ratios'):
config.load(root, 'test')
class SignalConfigTests(unittest.TestCase):
def test_signal_defaults(self):
item = SignalConfig()
self.assertEqual((item.url, item.timezone), ('', '*'))
self.assertFalse(item.gt_last_price_is_open)
self.assertEqual(GlobalConfig().signals, {})
def test_zero_hands_is_the_off_switch(self):
self.assertEqual(AccountConfig().zt_open_hands, 0)
if __name__ == '__main__':
unittest.main()