"""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()