optz
This commit is contained in:
147
labs/tests/test_etf_config.py
Normal file
147
labs/tests/test_etf_config.py
Normal file
@@ -0,0 +1,147 @@
|
||||
"""ETF 配置:固定文件加载、缺文件返回 None、全局默认与逐标的覆盖的校验。"""
|
||||
|
||||
import tempfile
|
||||
import textwrap
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import yaml
|
||||
|
||||
import config
|
||||
from config import EtfConfig, EtfDefaults, EtfSymbolConfig
|
||||
|
||||
GLOBAL = {"qmt_base_url": "unused", "api_host": "unused", "hosts": {"test": "account"}}
|
||||
SYMBOL = 'symbols: {"510300.SH": {is_t0: false, buy_shares: 1000, atr_multiplier: 1.0, inner_step: 0.7}}\n'
|
||||
|
||||
|
||||
class EtfConfigTests(unittest.TestCase):
|
||||
def load(self, etf: str | None = None, account: dict | None = None):
|
||||
"""在临时目录里生成 _global.yaml / account.yaml / _etf.yaml 并加载。"""
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "_global.yaml").write_text(
|
||||
yaml.safe_dump(dict(GLOBAL, qmt_data_dir=directory)), encoding="utf-8"
|
||||
)
|
||||
(root / "account.yaml").write_text(
|
||||
yaml.safe_dump(account or {"buy_value": 1000, "strategy": "etf"}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
if etf is not None:
|
||||
(root / "_etf.yaml").write_text(textwrap.dedent(etf), encoding="utf-8")
|
||||
with patch.object(config, "global_config"), \
|
||||
patch.object(config, "account_config"), \
|
||||
patch.object(config, "etf_config"):
|
||||
config.load(root, "test")
|
||||
return config.etf_config
|
||||
|
||||
def reject(self, etf: str, message: str):
|
||||
with self.assertRaisesRegex(ValueError, message):
|
||||
self.load(etf)
|
||||
|
||||
# ------------------------------------------------------------- 加载
|
||||
def test_missing_file_returns_none(self):
|
||||
"""_etf.yaml 缺失不是错误:其它策略的账户不应因此启动失败。"""
|
||||
self.assertIsNone(self.load())
|
||||
|
||||
def test_loads_defaults_and_symbols(self):
|
||||
loaded = self.load(
|
||||
"""
|
||||
defaults:
|
||||
atr_period: 10
|
||||
add_pct: 2.5
|
||||
symbols:
|
||||
"159915.SZ":
|
||||
is_t0: true
|
||||
buy_shares: 2000
|
||||
max_shares: 20000
|
||||
atr_multiplier: 1.5
|
||||
inner_step: 1.2
|
||||
"""
|
||||
)
|
||||
self.assertIsInstance(loaded, EtfConfig)
|
||||
self.assertEqual(loaded.codes, ("159915.SZ",))
|
||||
self.assertEqual(loaded.defaults.atr_period, 10)
|
||||
self.assertEqual(loaded.defaults.add_pct, 2.5)
|
||||
symbol = loaded.symbol("159915.SZ")
|
||||
self.assertTrue(symbol.is_t0)
|
||||
self.assertEqual(symbol.buy_shares, 2000)
|
||||
self.assertEqual(symbol.max_shares, 20000)
|
||||
|
||||
def test_unspecified_defaults_fall_back_to_dataclass(self):
|
||||
loaded = self.load(SYMBOL)
|
||||
self.assertEqual(loaded.defaults, EtfDefaults())
|
||||
self.assertEqual(loaded.defaults.commission_rate, 0.0003)
|
||||
|
||||
def test_symbol_inherits_uncovered_defaults(self):
|
||||
"""max_shares 缺省按 max_adds + 1 档计算,其余未覆盖项回落全局默认。"""
|
||||
loaded = self.load(SYMBOL)
|
||||
symbol = loaded.symbol("510300.SH")
|
||||
self.assertEqual(symbol.max_shares, 1000 * (loaded.defaults.max_adds + 1))
|
||||
self.assertEqual(symbol.get("rebound_pct"), loaded.defaults.rebound_pct)
|
||||
self.assertEqual(symbol.get("inner_grids"), loaded.defaults.inner_grids)
|
||||
self.assertEqual(symbol.get("max_hold_days"), loaded.defaults.max_hold_days)
|
||||
self.assertEqual(symbol.get("max_grid_span_pct"), loaded.defaults.max_grid_span_pct)
|
||||
|
||||
def test_symbol_override_wins_over_default(self):
|
||||
loaded = self.load(
|
||||
'defaults: {rebound_pct: 0.5}\n' + SYMBOL.replace(
|
||||
"inner_step: 0.7", "inner_step: 0.7, rebound_pct: 0.4"
|
||||
)
|
||||
)
|
||||
self.assertEqual(loaded.symbol("510300.SH").get("rebound_pct"), 0.4)
|
||||
|
||||
def test_unknown_symbol_is_rejected_by_lookup(self):
|
||||
loaded = self.load(SYMBOL)
|
||||
with self.assertRaisesRegex(KeyError, "159915.SZ"):
|
||||
loaded.symbol("159915.SZ")
|
||||
|
||||
# ------------------------------------------------------------- 校验
|
||||
def test_unknown_keys_are_rejected(self):
|
||||
self.reject("foo: 1\n", "foo")
|
||||
self.reject("defaults: {atr_period: 14, foo: 1}\n" + SYMBOL, "foo")
|
||||
self.reject(SYMBOL.replace("inner_step: 0.7", "inner_step: 0.7, add_pct: 2"), "add_pct")
|
||||
|
||||
def test_symbols_must_be_a_non_empty_mapping(self):
|
||||
self.reject("defaults: {atr_period: 14}\n", "symbols")
|
||||
|
||||
def test_symbol_code_must_be_a_listed_etf(self):
|
||||
self.reject('symbols: {"600000.SH": {is_t0: false, buy_shares: 100, '
|
||||
'atr_multiplier: 1.0, inner_step: 0.7}}\n', "600000.SH")
|
||||
|
||||
def test_required_symbol_fields_cannot_fall_back(self):
|
||||
"""is_t0、buy_shares、atr_multiplier、inner_step 没有安全的全局默认值。"""
|
||||
self.reject('symbols: {"510300.SH": {is_t0: false, buy_shares: 1000, '
|
||||
'atr_multiplier: 1.0}}\n', "inner_step")
|
||||
|
||||
def test_symbol_field_types_are_checked(self):
|
||||
self.reject(SYMBOL.replace("is_t0: false", "is_t0: 1"), "is_t0")
|
||||
self.reject(SYMBOL.replace("buy_shares: 1000", "buy_shares: 150"), "100")
|
||||
self.reject(SYMBOL.replace("buy_shares: 1000", "buy_shares: 0"), "buy_shares")
|
||||
self.reject(SYMBOL.replace("atr_multiplier: 1.0", "atr_multiplier: true"), "atr_multiplier")
|
||||
self.reject(SYMBOL.replace("atr_multiplier: 1.0", "atr_multiplier: 0"), "atr_multiplier")
|
||||
|
||||
def test_defaults_ranges_are_checked(self):
|
||||
self.reject("defaults: {max_adds: 10}\n" + SYMBOL, "max_adds")
|
||||
self.reject("defaults: {max_adds: -1}\n" + SYMBOL, "max_adds")
|
||||
self.reject("defaults: {max_adds: 1.5}\n" + SYMBOL, "max_adds")
|
||||
self.reject("defaults: {max_grid_span_pct: 101}\n" + SYMBOL, "max_grid_span_pct")
|
||||
self.reject("defaults: {atr_period: 0}\n" + SYMBOL, "atr_period")
|
||||
self.reject("defaults: {min_profit_pct: 0}\n" + SYMBOL, "min_profit_pct")
|
||||
# 0 表示不止损,是合法值;佣金允许为 0(回测口径)。
|
||||
loaded = self.load("defaults: {max_hold_days: 0, commission_rate: 0, min_commission: 0}\n" + SYMBOL)
|
||||
self.assertEqual(loaded.defaults.max_hold_days, 0)
|
||||
self.assertEqual(loaded.defaults.commission_rate, 0)
|
||||
|
||||
def test_rebound_must_be_below_add_pct(self):
|
||||
"""反弹确认价必须早于补仓触发,否则条件自相矛盾。"""
|
||||
self.reject("defaults: {rebound_pct: 3.0, add_pct: 3.0}\n" + SYMBOL, "rebound_pct")
|
||||
|
||||
def test_symbol_defaults_are_shared_not_copied(self):
|
||||
loaded = self.load(SYMBOL)
|
||||
self.assertIs(loaded.symbol("510300.SH").defaults, loaded.defaults)
|
||||
self.assertFalse(EtfSymbolConfig().is_t0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user