Files
big-qmt/labs/tests/test_etf_config.py
2026-09-19 19:45:43 +08:00

148 lines
6.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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()