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,324 @@
"""ETF 信号层:白名单展开、已收盘日线校验、指标计算与取数缓存。"""
from datetime import date, datetime, timedelta
from decimal import Decimal, ROUND_CEILING
from statistics import fmean
import unittest
from unittest.mock import Mock, patch
import httpx
from config import AccountConfig, EtfConfig, EtfDefaults, EtfSymbolConfig, GlobalConfig
from libs.signal import SignalItem
from strategy.etf import signal as etf_signal
CODE, OTHER = "510300.SH", "159915.SZ"
TODAY = date(2026, 9, 16)
# 最后一根日线为 2026-09-15前一交易日既不过期也不混入当日未收盘日线。
LAST_DAY = date(2026, 9, 15)
def raw_bar(day: str, close: float = 10.0, spread: float = 1.0, code: str = CODE) -> dict:
"""构造接口口径的一行日线:以 close 为中轴,上下各半个 spread。"""
close = float(close)
return dict(
ts_code=code,
trade_date=day,
open=close,
high=close + spread / 2,
low=close - spread / 2,
close=close,
)
def parse(*payload: dict, code: str = CODE, count: int = 120) -> list[dict]:
"""把接口口径的日线交给 parse_daily得到指标层使用的一行含 date"""
return etf_signal.parse_daily(list(payload), code, TODAY, count)
def raw_payload(
count: int = 90, close: float = 10.0, spread: float = 1.0, code: str = CODE
) -> list[dict]:
"""生成截至 LAST_DAY 的 count 根横盘日线(接口口径,升序)。"""
return [
raw_bar(
(LAST_DAY - timedelta(days=count - 1 - index)).strftime("%Y%m%d"),
close,
spread,
code,
)
for index in range(count)
]
def bars(count: int = 90, close: float = 10.0, spread: float = 1.0) -> list[dict]:
"""生成指标层口径(已解析)的 count 根横盘日线。"""
return parse(*raw_payload(count, close, spread))
def symbol(**overrides) -> EtfSymbolConfig:
base = dict(is_t0=False, buy_shares=1000, atr_multiplier=1.0, inner_step=0.7)
base.update(overrides)
return EtfSymbolConfig(**{k: v for k, v in base.items() if v is not None})
def runtime(etf: EtfConfig | None = None, host: str = "http://api.test") -> Mock:
rt = Mock()
rt.etf_cfg = etf if etf is not None else EtfConfig(
defaults=EtfDefaults(), symbols={CODE: symbol()}
)
rt.global_cfg = GlobalConfig(api_host=host)
rt.account_cfg = AccountConfig(strategy="etf")
return rt
class ParseDailyTests(unittest.TestCase):
def test_bare_list_is_supported_and_sorted_ascending(self):
payload = [raw_bar("20260915"), raw_bar("20260912")]
rows = etf_signal.parse_daily(payload, CODE, TODAY)
self.assertEqual([row["date"] for row in rows], ["20260912", "20260915"])
self.assertEqual(rows[0]["close"], 10.0)
def test_legacy_envelope_is_supported(self):
payload = {"code": 0, "message": "", "details": [raw_bar("20260915")]}
self.assertEqual(
[row["date"] for row in etf_signal.parse_daily(payload, CODE, TODAY)],
["20260915"],
)
def test_today_and_future_bars_are_dropped_then_limited(self):
payload = [
raw_bar("20260916"), raw_bar("20260917"),
raw_bar("20260915"), raw_bar("20260914"),
]
rows = etf_signal.parse_daily(payload, CODE, TODAY, count=1)
self.assertEqual([row["date"] for row in rows], ["20260915"])
def test_numeric_strings_are_converted(self):
rows = etf_signal.parse_daily([raw_bar("20260915", close="10.5")], CODE, TODAY)
self.assertEqual(rows[0]["close"], 10.5)
def test_bad_payloads_are_rejected(self):
cases = [
None, [], {}, {"code": 1, "message": "failed"},
{"code": False, "details": [raw_bar("20260915")]},
{"code": 0, "details": []}, {"code": 0, "details": None},
]
for payload in cases:
with self.subTest(payload=payload), self.assertRaises(ValueError):
etf_signal.parse_daily(payload, CODE, TODAY)
def test_bad_rows_are_rejected(self):
cases = [
[raw_bar("20260915", code=OTHER)], # 证券归属不一致
[raw_bar("20260915"), raw_bar("20260915")], # 日期重复
[raw_bar("20260230")], # 非法日期
[raw_bar("20260915", close=float("nan"))], # 非有限
[dict(raw_bar("20260915"), open=True)], # 布尔价格
[dict(raw_bar("20260915"), low=12)], # OHLC 关系异常
[dict(raw_bar("20260915"), close=None)], # 缺失价格
]
for rows in cases:
with self.subTest(rows=rows), self.assertRaises(ValueError):
etf_signal.parse_daily(rows, CODE, TODAY)
def test_count_must_be_positive_integer(self):
for count in (0, -1, 1.5, True):
with self.subTest(count=count), self.assertRaises(ValueError):
etf_signal.parse_daily([raw_bar("20260915")], CODE, TODAY, count)
class CalculateTests(unittest.TestCase):
"""横盘日线ATR = spread、MA60 = close、格距 = max(ATR×倍数, MA60×0.5%, 0.001)。"""
def test_flat_bars_produce_expected_indicators(self):
values = etf_signal.calculate(bars(), symbol(), EtfDefaults(), TODAY)
self.assertEqual(values[etf_signal.IND_MA60], 10.0)
self.assertEqual(values[etf_signal.IND_ATR], 1.0)
self.assertEqual(values[etf_signal.IND_CHANNEL_LOW], 9.5)
self.assertEqual(values[etf_signal.IND_CHANNEL_HIGH], 10.5)
# 入场门槛 = min(9.5 + 1×15%, 10) = 9.65
self.assertAlmostEqual(values[etf_signal.IND_ENTRY], 9.65)
self.assertEqual(values[etf_signal.IND_GRID], 1.0)
self.assertEqual(values[etf_signal.IND_PRICE], 10.0)
self.assertAlmostEqual(values[etf_signal.IND_GRID_PCT], 10.0)
self.assertAlmostEqual(values[etf_signal.IND_ADD_PRICE], 9.7)
def test_atr_multiplier_and_grid_floor(self):
low_atr = etf_signal.calculate(bars(spread=0.02), symbol(atr_multiplier=1.0),
EtfDefaults(), TODAY)
# ATR=0.02 低于 MA60×0.5% = 0.05 的百分比下限
self.assertAlmostEqual(low_atr[etf_signal.IND_GRID], 0.05)
wide = etf_signal.calculate(bars(), symbol(atr_multiplier=0.5), EtfDefaults(), TODAY)
self.assertEqual(wide[etf_signal.IND_GRID], 0.5)
def test_grid_is_rounded_up_to_tick(self):
rows = bars()
for row in rows:
row["close"] += 0.0004
row["open"], row["high"], row["low"] = row["close"], row["high"] + 0.0004, row["low"] + 0.0004
values = etf_signal.calculate(rows, symbol(), EtfDefaults(), TODAY)
raw = values[etf_signal.IND_ATR] * 1.0
expected = float(Decimal(str(raw)).quantize(Decimal("0.001"), rounding=ROUND_CEILING))
self.assertEqual(values[etf_signal.IND_GRID], expected)
self.assertEqual(round(values[etf_signal.IND_GRID], 3), values[etf_signal.IND_GRID])
def test_rising_close_uses_wilder_smoothing(self):
rows = [dict(row) for row in bars()]
for index, row in enumerate(rows):
shift = index * 0.01
row["close"] += shift
row["open"], row["high"], row["low"] = (
row["close"], row["close"] + 0.5, row["close"] - 0.5
)
values = etf_signal.calculate(rows, symbol(), EtfDefaults(), TODAY)
closes = [row["close"] for row in rows]
self.assertAlmostEqual(values[etf_signal.IND_MA60], fmean(closes[-60:]))
self.assertGreaterEqual(values[etf_signal.IND_ATR], 1.0)
def test_insufficient_or_stale_bars_are_rejected(self):
with self.assertRaisesRegex(ValueError, "61"):
etf_signal.calculate(bars(40), symbol(), EtfDefaults(), TODAY)
# 最近日线距今超过 15 个自然日即放弃该标的
with self.assertRaisesRegex(ValueError, "自然日"):
etf_signal.calculate(bars(), symbol(), EtfDefaults(), TODAY + timedelta(days=20))
def test_bad_atr_period_is_rejected(self):
defaults = EtfDefaults()
defaults.atr_period = 1
with self.assertRaisesRegex(ValueError, "atr_period"):
etf_signal.calculate(bars(), symbol(), defaults, TODAY)
class DailyBarsTests(unittest.TestCase):
def test_request_url_and_no_token_header(self):
def respond(request):
self.assertEqual(
str(request.url), etf_signal.DAILY_URL + "?code=" + CODE
)
self.assertNotIn("x-token", request.headers)
return httpx.Response(200, json=[raw_bar("20260915")])
with httpx.Client(transport=httpx.MockTransport(respond)) as client:
rows = etf_signal.daily_bars(client, CODE, TODAY)
self.assertEqual([row["date"] for row in rows], ["20260915"])
def test_custom_endpoint_is_used(self):
def respond(request):
self.assertEqual(str(request.url), "http://api.test/etf/daily?code=" + CODE)
return httpx.Response(200, json=[raw_bar("20260915")])
with httpx.Client(transport=httpx.MockTransport(respond)) as client:
etf_signal.daily_bars(
client, CODE, TODAY, endpoint="http://api.test/etf/daily"
)
def test_http_and_json_errors_propagate(self):
for status, content in ((404, "{}"), (200, "<html>error</html>")):
with httpx.Client(
transport=httpx.MockTransport(lambda r: httpx.Response(status, text=content))
) as client, self.assertRaises((httpx.HTTPStatusError, ValueError)):
etf_signal.daily_bars(client, CODE, TODAY)
def daily_response(payload: object, code: str = CODE) -> httpx.Response:
"""构造带 request 的 200 响应httpx 的 raise_for_status 需要 request。"""
request = httpx.Request("GET", etf_signal.DAILY_URL, params={"code": code})
return httpx.Response(200, json=payload, request=request)
def boom(url=None, params=None):
"""模拟连接失败httpx 只接受真实的 RequestError 子类实例。"""
raise httpx.ConnectError("boom")
class GenSignalsTests(unittest.TestCase):
def setUp(self):
self.client = Mock(spec=httpx.Client)
self.client.get.side_effect = lambda url, params=None: daily_response(
raw_payload(code=params["code"]), params["code"]
)
for cache in (etf_signal._daily_cache, etf_signal._fetched, etf_signal._retry_at):
cache.clear()
patch.object(etf_signal, "_history_client", self.client).start()
self.addCleanup(patch.stopall)
def config(self, codes=(CODE,)):
return EtfConfig(
defaults=EtfDefaults(),
symbols={code: symbol() for code in codes},
)
def test_signals_follow_config_order_and_carry_indicators(self):
signals = etf_signal.gen_signals(runtime(self.config((OTHER, CODE))))
self.assertEqual([item.code for item in signals], [OTHER, CODE])
item = signals[0]
self.assertIsInstance(item, SignalItem)
self.assertEqual(item.signal_key, "etf")
self.assertEqual(item.last_close, 10.0)
self.assertIn("ETF网格", item.desc)
self.assertEqual(item.tech_indicator[etf_signal.IND_GRID], 1.0)
self.assertEqual(item.tech_indicator[etf_signal.IND_MA60], 10.0)
self.assertEqual(item.tech_indicator[etf_signal.IND_ENTRY], 9.65)
def test_missing_etf_config_yields_empty_list(self):
rt = runtime()
rt.etf_cfg = None
self.assertEqual(etf_signal.gen_signals(rt), [])
def test_broken_symbol_is_skipped_without_breaking_others(self):
"""CODE 的日线证券代码不一致时放弃该标的OTHER 正常生成。"""
payloads = {CODE: [raw_bar("20260915", code=OTHER)], OTHER: raw_payload(code=OTHER)}
self.client.get.side_effect = lambda url, params=None: daily_response(
payloads[params["code"]], params["code"]
)
signals = etf_signal.gen_signals(runtime(self.config((CODE, OTHER))))
self.assertEqual([item.code for item in signals], [OTHER])
def test_daily_bars_are_cached_per_code_per_day(self):
rt = runtime(self.config())
etf_signal.gen_signals(rt)
etf_signal.gen_signals(rt)
self.assertEqual(self.client.get.call_count, 1)
def test_failure_is_throttled_by_retry_window(self):
self.client.get.side_effect = boom
rt = runtime(self.config())
self.assertEqual(etf_signal.gen_signals(rt), [])
self.assertEqual(self.client.get.call_count, 1)
# 重试窗口内不再取数
self.assertEqual(etf_signal.gen_signals(rt), [])
self.assertEqual(self.client.get.call_count, 1)
self.assertIn(CODE, etf_signal._retry_at)
def test_retry_happens_after_window(self):
self.client.get.side_effect = boom
rt = runtime(self.config())
etf_signal.gen_signals(rt)
etf_signal._retry_at[CODE] = datetime.now() - timedelta(seconds=1)
self.client.get.side_effect = None
self.client.get.return_value = daily_response(raw_payload())
signals = etf_signal.gen_signals(rt)
self.assertEqual([item.code for item in signals], [CODE])
self.assertEqual(self.client.get.call_count, 2)
def test_new_trading_day_clears_cache(self):
rt = runtime(self.config())
etf_signal.gen_signals(rt)
# 模拟隔日:昨日的取数记录不应继续复用。
etf_signal._fetched[CODE] = LAST_DAY - timedelta(days=1)
etf_signal._daily_cache.clear()
etf_signal.gen_signals(rt)
self.assertEqual(etf_signal._fetched[CODE], datetime.now().date())
self.assertEqual(self.client.get.call_count, 2)
def test_endpoint_falls_back_to_default_without_api_host(self):
etf_signal.gen_signals(runtime(self.config(), host=""))
self.assertEqual(self.client.get.call_args.args[0], etf_signal.DAILY_URL)
def test_endpoint_uses_global_api_host(self):
etf_signal.gen_signals(runtime(self.config(), host="http://api.test/"))
self.assertEqual(self.client.get.call_args.args[0], "http://api.test/etf/daily")
if __name__ == "__main__":
unittest.main()