325 lines
14 KiB
Python
325 lines
14 KiB
Python
"""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()
|