This commit is contained in:
2026-09-19 19:45:43 +08:00
parent 7183cb45f8
commit 8131b158b4
60 changed files with 7669 additions and 909 deletions

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