"""Point-in-time earnings calendar helpers for leakage-safe future event scheduling.""" from __future__ import annotations import datetime as dt from bisect import bisect_right from dataclasses import dataclass from functools import lru_cache from pathlib import Path from typing import Any, Iterable import requests import pyarrow.parquet as pq from libs.common.logging import get_logger from libs.labeler.reaction_date import compute_reaction_date logger = get_logger(__name__) @dataclass(frozen=True) class EarningsCalendarEntry: symbol: str as_of_date: dt.date expected_reaction_date: dt.date expected_event_date: dt.date | None = None filing_time_bucket: str | None = None confidence: float | None = None revision_count: int | None = None is_cancelled: bool = False source: str | None = None def _coerce_date(value: Any) -> dt.date | None: if value is None: return None if isinstance(value, dt.datetime): return value.date() if isinstance(value, dt.date): return value text = str(value).strip() if not text: return None try: return dt.date.fromisoformat(text[:10]) except ValueError: return None def _coerce_float(value: Any) -> float | None: if value is None: return None try: return float(value) except (TypeError, ValueError): return None def _coerce_int(value: Any) -> int | None: if value is None: return None try: return int(value) except (TypeError, ValueError): return None def _coerce_bool(value: Any) -> bool: if isinstance(value, bool): return value if value is None: return False if isinstance(value, (int, float)): return bool(value) text = str(value).strip().lower() return text in {"1", "true", "t", "yes", "y"} class PointInTimeEarningsCalendar: """Latest-known earnings schedule for each symbol as of a prior date.""" def __init__(self, entries: Iterable[EarningsCalendarEntry]) -> None: grouped: dict[str, list[EarningsCalendarEntry]] = {} for entry in entries: grouped.setdefault(entry.symbol, []).append(entry) self._entries_by_symbol: dict[str, tuple[list[dt.date], list[EarningsCalendarEntry]]] = {} for symbol, symbol_entries in grouped.items(): symbol_entries.sort(key=lambda entry: (entry.as_of_date, entry.expected_reaction_date)) as_of_dates = [entry.as_of_date for entry in symbol_entries] self._entries_by_symbol[symbol] = (as_of_dates, symbol_entries) @classmethod def from_parquet(cls, path: Path) -> PointInTimeEarningsCalendar: table = pq.read_table(str(path)) rows = table.to_pylist() entries: list[EarningsCalendarEntry] = [] for row in rows: symbol = str( row.get("symbol") or row.get("ticker") or row.get("trade_symbol") or "" ).strip().upper() if not symbol: continue as_of_date = _coerce_date( row.get("as_of_date") or row.get("known_as_of_date") or row.get("snapshot_date") or row.get("published_date") ) if as_of_date is None: continue expected_event_date = _coerce_date( row.get("expected_event_date") or row.get("event_date") or row.get("earnings_date") or row.get("next_earnings_date") ) filing_time_bucket = ( row.get("expected_filing_time_bucket") or row.get("filing_time_bucket") or row.get("timing_class") ) expected_reaction_date = _coerce_date( row.get("expected_reaction_date") or row.get("reaction_date") or row.get("expected_execution_date") ) if expected_reaction_date is None and expected_event_date is not None: expected_reaction_date = compute_reaction_date( expected_event_date, str(filing_time_bucket or "post_market"), ) if expected_reaction_date is None: continue entries.append( EarningsCalendarEntry( symbol=symbol, as_of_date=as_of_date, expected_reaction_date=expected_reaction_date, expected_event_date=expected_event_date, filing_time_bucket=str(filing_time_bucket) if filing_time_bucket is not None else None, confidence=_coerce_float(row.get("confidence") or row.get("resolver_confidence")), revision_count=_coerce_int(row.get("revision_count")), is_cancelled=_coerce_bool(row.get("is_cancelled") or row.get("cancelled")), source=str(row.get("source")) if row.get("source") is not None else None, ) ) logger.info( "pit_earnings_calendar_loaded", path=str(path), rows=len(entries), symbols=len({entry.symbol for entry in entries}), ) return cls(entries) def get_known_upcoming_reaction_dates( self, as_of_date: dt.date, allowed_reaction_dates: Iterable[dt.date], symbols: Iterable[str] | None = None, ) -> dict[str, dt.date]: allowed = set(allowed_reaction_dates) if not allowed: return {} symbol_filter = { str(symbol).strip().upper() for symbol in (symbols or []) if str(symbol).strip() } result: dict[str, dt.date] = {} for symbol, (as_of_dates, entries) in self._entries_by_symbol.items(): if symbol_filter and symbol not in symbol_filter: continue idx = bisect_right(as_of_dates, as_of_date) - 1 if idx < 0: continue entry = entries[idx] if entry.is_cancelled: continue if entry.expected_reaction_date in allowed: result[symbol] = entry.expected_reaction_date return result def _normalize_earnings_time_bucket(value: Any) -> str: text = str(value or "unknown").strip().lower() if text == "during_market": return "regular_hours" if text in {"pre_market", "regular_hours", "post_market", "unknown"}: return text return "unknown" class OraclePointInTimeEarningsCalendar: """Oracle-backed PIT upcoming earnings lookup using Stock Oracle bulk API.""" def __init__( self, oracle_url: str, timeout: float = 30.0, session: requests.Session | None = None, ) -> None: self._base_url = oracle_url.rstrip("/") if oracle_url else "" self._timeout = float(timeout) self._session = session or requests.Session() self._cache: dict[tuple[str, tuple[str, ...], tuple[str, ...]], dict[str, dt.date]] = {} def get_known_upcoming_reaction_dates( self, as_of_date: dt.date, allowed_reaction_dates: Iterable[dt.date], symbols: Iterable[str] | None = None, ) -> dict[str, dt.date]: if not self._base_url: return {} allowed = sorted(set(allowed_reaction_dates)) requested_symbols = sorted( { str(symbol).strip().upper() for symbol in (symbols or []) if str(symbol).strip() } ) if not allowed or not requested_symbols: return {} max_days_ahead = max((reaction_date - as_of_date).days for reaction_date in allowed) if max_days_ahead < 1: return {} allowed_key = tuple(date.isoformat() for date in allowed) symbol_key = tuple(requested_symbols) cache_key = (as_of_date.isoformat(), allowed_key, symbol_key) cached = self._cache.get(cache_key) if cached is not None: return dict(cached) results: dict[str, dt.date] = {} limit = max(1, min(20, len(allowed) + 1)) for idx in range(0, len(requested_symbols), 50): batch = requested_symbols[idx: idx + 50] payload = { "symbols": batch, "days_ahead": min(365, max_days_ahead), "limit": limit, "as_of_date": as_of_date.isoformat(), } try: response = self._session.post( f"{self._base_url}/api/v1/earnings/calendar/bulk", json=payload, timeout=self._timeout, ) response.raise_for_status() body = response.json() except Exception as exc: logger.warning( "oracle_pit_earnings_calendar_fetch_failed", error=str(exc), as_of_date=as_of_date.isoformat(), symbol_count=len(batch), ) return {} for entry in body.get("entries", []): symbol = str(entry.get("symbol") or "").strip().upper() if not symbol or symbol in results: continue earnings_date = _coerce_date(entry.get("earnings_date")) if earnings_date is None: continue reaction_date = compute_reaction_date( earnings_date, _normalize_earnings_time_bucket(entry.get("earnings_time")), ) if reaction_date not in allowed: continue results[symbol] = reaction_date logger.info( "oracle_pit_earnings_calendar_loaded", as_of_date=as_of_date.isoformat(), symbol_count=len(requested_symbols), matches=len(results), ) self._cache[cache_key] = dict(results) return results @lru_cache(maxsize=8) def load_pit_earnings_calendar(path_str: str) -> PointInTimeEarningsCalendar | None: path = Path(path_str) if not path.exists(): logger.info("pit_earnings_calendar_missing", path=str(path)) return None return PointInTimeEarningsCalendar.from_parquet(path) __all__ = [ "EarningsCalendarEntry", "OraclePointInTimeEarningsCalendar", "PointInTimeEarningsCalendar", "load_pit_earnings_calendar", ]