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
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
|