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.

253 lines
10 KiB
Python

"""Attention filtering service — shared by BacktestRunner and PaperTradingEngine.
Fetches event attention data (Wikipedia pageviews, GDELT news, entity resolver)
from Stock Oracle and applies engine-level gates + rescoring.
"""
from __future__ import annotations
import datetime as dt
from typing import Any
import requests
from libs.backtest.domain import Candidate, SignalConfig
from libs.backtest.selector import rank_candidates
from libs.common.logging import get_logger
logger = get_logger(__name__)
# Lazy import to avoid circular deps
EventAttentionResponse = None
def _get_attention_response_class():
global EventAttentionResponse
if EventAttentionResponse is None:
from libs.oracle_client.models import EventAttentionResponse as _cls
EventAttentionResponse = _cls
return EventAttentionResponse
class AttentionFilterService:
"""Fetches and caches event attention data, applies gates and rescoring."""
def __init__(self, oracle_url: str, scoring_model: str) -> None:
self._base_url = oracle_url.rstrip("/") if oracle_url else ""
self._scoring_model = scoring_model
self._cache: dict[tuple[str, dt.date | None], Any] = {}
self._session: requests.Session | None = None
if self._base_url:
self._session = requests.Session()
def close(self) -> None:
if self._session:
self._session.close()
self._session = None
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def engine_requires_attention(self, engine: Any) -> bool:
"""Return True if this engine needs attention data for filtering/rescoring."""
_ATTENTION_MODELS = {
"return_max_long_v1", "return_max_long_v2", "return_max_long_v3",
"return_max_long_v4", "return_max_long_v5", "return_max_long_v6",
"return_max_long_v7", "return_max_long_v8", "return_max_long_v9",
"return_max_long_v9g", "return_max_long_v10",
}
if self._scoring_model in _ATTENTION_MODELS and engine.direction != "short_only":
return True
return any(
value is not None
for value in (
engine.attention_min_wiki_spike_10d,
engine.attention_min_wiki_zscore_20d,
engine.attention_max_wiki_spike_10d,
engine.attention_max_wiki_zscore_20d,
engine.attention_min_article_count_3d,
engine.attention_min_us_article_count_3d,
engine.attention_min_resolver_confidence,
)
)
def engine_requires_attention_data(self, engine: Any) -> bool:
"""Return True when the engine has minimum-style gates (need actual data)."""
return any(
value is not None
for value in (
engine.attention_min_wiki_spike_10d,
engine.attention_min_wiki_zscore_20d,
engine.attention_min_article_count_3d,
engine.attention_min_us_article_count_3d,
engine.attention_min_resolver_confidence,
)
)
def apply_filters(
self,
candidates: list[Candidate],
engine: Any,
signal_config: SignalConfig,
) -> list[Candidate]:
"""Apply attention gates and rescoring — matches BacktestRunner._apply_attention_filters."""
if not candidates or not self.engine_requires_attention(engine):
return candidates[: signal_config.max_candidates_per_day]
filtered: list[Candidate] = []
requires_data = self.engine_requires_attention_data(engine)
threshold = (
engine.score_threshold_override
if engine.score_threshold_override is not None
else signal_config.score_threshold
)
for candidate in candidates:
attention = self._get_event_attention(candidate)
if attention is None:
if not requires_data:
filtered.append(candidate)
continue
if not self._passes_filters(engine, attention):
continue
enriched = self._attach_features(candidate, attention)
enriched = self._maybe_rescore(enriched)
if enriched.score >= threshold:
filtered.append(enriched)
filtered = rank_candidates(filtered, signal_config.ranking_fields)
return filtered[: signal_config.max_candidates_per_day]
# ------------------------------------------------------------------
# Internal
# ------------------------------------------------------------------
def _get_event_attention(self, candidate: Candidate) -> Any:
event_date = candidate.event_date or candidate.reaction_date
cache_key = (candidate.symbol, event_date)
if cache_key in self._cache:
return self._cache[cache_key]
if not self._base_url or self._session is None:
self._cache[cache_key] = None
return None
cls = _get_attention_response_class()
try:
response = self._session.get(
f"{self._base_url}/api/v1/attention/event/{candidate.symbol}",
params={"event_date": event_date.isoformat() if event_date else ""},
timeout=30,
)
if response.status_code >= 400:
logger.debug(
"attention_fetch_failed",
symbol=candidate.symbol,
status_code=response.status_code,
)
self._cache[cache_key] = None
return None
payload = cls.model_validate(response.json())
except Exception as exc:
logger.warning(
"attention_fetch_error",
symbol=candidate.symbol,
error=str(exc),
)
self._cache[cache_key] = None
return None
self._cache[cache_key] = payload
return payload
def _passes_filters(self, engine: Any, attention: Any) -> bool:
if (
engine.attention_min_wiki_spike_10d is not None
and (attention.wiki.spike_10d is None or attention.wiki.spike_10d < engine.attention_min_wiki_spike_10d)
):
return False
if (
engine.attention_min_wiki_zscore_20d is not None
and (attention.wiki.zscore_20d is None or attention.wiki.zscore_20d < engine.attention_min_wiki_zscore_20d)
):
return False
if (
engine.attention_max_wiki_spike_10d is not None
and attention.wiki.spike_10d is not None
and attention.wiki.spike_10d > engine.attention_max_wiki_spike_10d
):
return False
if (
engine.attention_max_wiki_zscore_20d is not None
and attention.wiki.zscore_20d is not None
and attention.wiki.zscore_20d > engine.attention_max_wiki_zscore_20d
):
return False
if (
engine.attention_min_article_count_3d is not None
and attention.news.article_count_3d < engine.attention_min_article_count_3d
):
return False
if (
engine.attention_min_us_article_count_3d is not None
and attention.news.us_article_count_3d < engine.attention_min_us_article_count_3d
):
return False
if (
engine.attention_min_resolver_confidence is not None
and attention.entity.resolver_confidence < engine.attention_min_resolver_confidence
):
return False
return True
def _attach_features(self, candidate: Candidate, attention: Any) -> Candidate:
features = dict(candidate.features)
features.update({
"attention_wiki_spike_10d": attention.wiki.spike_10d,
"attention_wiki_zscore_20d": attention.wiki.zscore_20d,
"attention_article_count_3d": attention.news.article_count_3d,
"attention_us_article_count_3d": attention.news.us_article_count_3d,
"attention_gdelt_status": attention.news.gdelt_status,
"attention_resolver_confidence": attention.entity.resolver_confidence,
})
return candidate.model_copy(update={"features": features})
def _maybe_rescore(self, candidate: Candidate) -> Candidate:
_RESCORE_MODELS = {
"return_max_long_v1", "return_max_long_v2", "return_max_long_v3",
"return_max_long_v4", "return_max_long_v5", "return_max_long_v6",
"return_max_long_v7", "return_max_long_v8",
}
if self._scoring_model not in _RESCORE_MODELS:
return candidate
from libs.backtest.scoring import (
compute_return_max_long_score,
compute_return_max_long_score_v2,
compute_return_max_long_score_v3,
compute_return_max_long_score_v4,
compute_return_max_long_score_v5,
compute_return_max_long_score_v6,
compute_return_max_long_score_v7,
compute_return_max_long_score_v8,
)
rescored_features = dict(candidate.features)
rescored_features.update({
"event_type": candidate.event_type,
"event_direction": rescored_features.get("event_direction"),
})
_dispatch = {
"return_max_long_v8": compute_return_max_long_score_v8,
"return_max_long_v7": compute_return_max_long_score_v7,
"return_max_long_v6": compute_return_max_long_score_v6,
"return_max_long_v5": compute_return_max_long_score_v5,
"return_max_long_v4": compute_return_max_long_score_v4,
"return_max_long_v3": compute_return_max_long_score_v3,
"return_max_long_v2": compute_return_max_long_score_v2,
"return_max_long_v1": compute_return_max_long_score,
}
scoring_fn = _dispatch.get(self._scoring_model, compute_return_max_long_score)
score = scoring_fn(rescored_features)
return candidate.model_copy(update={"score": score})