You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

101 lines
3.1 KiB
Python

from __future__ import annotations
import datetime as dt
from types import SimpleNamespace
from zoneinfo import ZoneInfo
import requests
from libs.backtest.attention import AttentionFilterService
from libs.backtest.domain import Candidate, SignalConfig
_UTC = ZoneInfo("UTC")
def _make_candidate(symbol: str, event_id: str) -> Candidate:
return Candidate(
event_id=event_id,
symbol=symbol,
score=0.8,
sector="Technology",
event_type="earnings_release",
event_timestamp=dt.datetime(2026, 1, 5, 21, 0, tzinfo=_UTC),
event_date=dt.date(2026, 1, 5),
filing_time_bucket="post_market",
reaction_date=dt.date(2026, 1, 6),
execution_date=dt.date(2026, 1, 7),
entry_price_est=100.0,
avg_dollar_volume=10_000_000.0,
score_bucket="high",
)
def _make_engine(**overrides):
defaults = {
"direction": "long_only",
"attention_min_wiki_spike_10d": None,
"attention_min_wiki_zscore_20d": None,
"attention_max_wiki_spike_10d": None,
"attention_max_wiki_zscore_20d": None,
"attention_min_article_count_3d": None,
"attention_min_us_article_count_3d": None,
"attention_min_resolver_confidence": None,
"score_threshold_override": None,
}
defaults.update(overrides)
return SimpleNamespace(**defaults)
class _TimeoutSession:
def __init__(self) -> None:
self.calls = 0
self.timeouts: list[float] = []
def get(self, *_args, **kwargs):
self.calls += 1
self.timeouts.append(float(kwargs["timeout"]))
raise requests.ReadTimeout("oracle unavailable")
def close(self) -> None:
return None
def test_attention_service_disables_run_after_network_timeout_for_max_only_gates():
service = AttentionFilterService("http://oracle:18001", "default", timeout=1.5)
session = _TimeoutSession()
service._session = session
engine = _make_engine(attention_max_wiki_spike_10d=6.0)
signal = SignalConfig(score_threshold=0.5, max_candidates_per_day=10)
candidates = [
_make_candidate("AAPL", "EVT::1"),
_make_candidate("MSFT", "EVT::2"),
]
filtered = service.apply_filters(candidates, engine, signal)
assert [candidate.symbol for candidate in filtered] == ["AAPL", "MSFT"]
assert session.calls == 1
assert session.timeouts == [1.5]
assert service._service_unavailable is True
def test_attention_service_disables_run_after_network_timeout_for_min_gates():
service = AttentionFilterService("http://oracle:18001", "default", timeout=2.0)
session = _TimeoutSession()
service._session = session
engine = _make_engine(attention_min_wiki_spike_10d=1.0)
signal = SignalConfig(score_threshold=0.5, max_candidates_per_day=10)
candidates = [
_make_candidate("AAPL", "EVT::1"),
_make_candidate("MSFT", "EVT::2"),
]
filtered = service.apply_filters(candidates, engine, signal)
assert filtered == []
assert session.calls == 1
assert session.timeouts == [2.0]
assert service._service_unavailable is True