""" Session-aggregated news/social signal. Computes per-(ticker, session_date, window) aggregates directly from raw NewsHeadline rows via SQL — no materialized table. Redis caching at the endpoint layer absorbs repeated reads. Aggregate fields: headline_count, primary_count, first/last_headline_at, category_counts (JSON), sentiment_mean, sentiment_recency_weighted, social: {message_count, bull_count, bear_count, bull_bear_ratio}, sources_present """ from __future__ import annotations import logging import math from collections import Counter from dataclasses import asdict, dataclass, field from datetime import date, datetime, timezone from typing import Iterable from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.news_headline import NewsHeadline from app.services.news.session_window import WindowName, session_window logger = logging.getLogger(__name__) # Recency-weight half-life: 1 hour close to open contributes weight 1.0, # 6 hours earlier ≈ 0.55. Matches plan (decay=0.1 per hour). _RECENCY_DECAY_PER_MIN = 0.1 / 60.0 @dataclass class SocialStats: message_count: int = 0 bull_count: int = 0 bear_count: int = 0 bull_bear_ratio: float | None = None # bull/(bull+bear); None if no directional msgs @dataclass class SessionAggregate: ticker: str session_date: str # ISO date (ET) window: str headline_count: int = 0 primary_count: int = 0 first_headline_at: str | None = None last_headline_at: str | None = None category_counts: dict[str, int] = field(default_factory=dict) sentiment_mean: float | None = None sentiment_recency_weighted: float | None = None social: SocialStats = field(default_factory=SocialStats) sources_present: list[str] = field(default_factory=list) def to_dict(self) -> dict: return asdict(self) async def aggregate_session( db: AsyncSession, ticker: str, session_date: date, window: WindowName, sources: list[str] | None = None, pit_cutoff: datetime | None = None, ) -> SessionAggregate: """ Aggregate news headlines for one (ticker, session_date, window). `pit_cutoff` (optional) excludes rows whose `ingested_at > cutoff` to enable lookahead-safe backtest queries. If omitted, defaults to the session window's end_utc (so we never see headlines that arrived after the window closed in real time). """ ticker_u = ticker.strip().upper() start_utc, end_utc = session_window(session_date, window) cutoff = pit_cutoff or end_utc stmt = select(NewsHeadline).where( NewsHeadline.ticker == ticker_u, NewsHeadline.published_at >= start_utc, NewsHeadline.published_at < end_utc, NewsHeadline.ingested_at <= cutoff, ) if sources: stmt = stmt.where(NewsHeadline.source.in_([s.strip() for s in sources if s])) result = await db.execute(stmt) rows = result.scalars().all() return _build_aggregate( rows=rows, ticker=ticker_u, session_date=session_date, window=window, window_end_utc=end_utc, ) async def aggregate_session_batch( db: AsyncSession, tickers: list[str], session_date: date, window: WindowName, sources: list[str] | None = None, pit_cutoff: datetime | None = None, ) -> dict[str, SessionAggregate]: """Aggregate for many tickers in one query. Missing tickers → empty agg.""" tickers_u = [t.strip().upper() for t in tickers if t and t.strip()] if not tickers_u: return {} start_utc, end_utc = session_window(session_date, window) cutoff = pit_cutoff or end_utc stmt = select(NewsHeadline).where( NewsHeadline.ticker.in_(tickers_u), NewsHeadline.published_at >= start_utc, NewsHeadline.published_at < end_utc, NewsHeadline.ingested_at <= cutoff, ) if sources: stmt = stmt.where(NewsHeadline.source.in_([s.strip() for s in sources if s])) result = await db.execute(stmt) rows = result.scalars().all() by_ticker: dict[str, list[NewsHeadline]] = {t: [] for t in tickers_u} for r in rows: by_ticker.setdefault(r.ticker, []).append(r) out: dict[str, SessionAggregate] = {} for t in tickers_u: out[t] = _build_aggregate( rows=by_ticker.get(t, []), ticker=t, session_date=session_date, window=window, window_end_utc=end_utc, ) return out # --------------------------------------------------------------------------- # Aggregate construction # --------------------------------------------------------------------------- def _build_aggregate( rows: Iterable[NewsHeadline], ticker: str, session_date: date, window: str, window_end_utc: datetime, ) -> SessionAggregate: rows = list(rows) agg = SessionAggregate( ticker=ticker, session_date=session_date.isoformat(), window=window, ) if not rows: return agg cat_counter: Counter[str] = Counter() sentiments: list[float] = [] weighted_num = 0.0 weighted_den = 0.0 sources: set[str] = set() social_msgs = 0 social_bull = 0 social_bear = 0 first_at: datetime | None = None last_at: datetime | None = None primary_count = 0 for r in rows: sources.add(r.source) if r.is_primary: primary_count += 1 for c in (r.categories or []): cat_counter[c] += 1 pub = r.published_at if pub.tzinfo is None: pub = pub.replace(tzinfo=timezone.utc) if first_at is None or pub < first_at: first_at = pub if last_at is None or pub > last_at: last_at = pub if r.raw_sentiment is not None: sentiments.append(float(r.raw_sentiment)) # Recency weight: closer to window end_utc → higher weight mins_before_end = max(0.0, (window_end_utc - pub).total_seconds() / 60.0) w = math.exp(-_RECENCY_DECAY_PER_MIN * mins_before_end) weighted_num += w * float(r.raw_sentiment) weighted_den += w if r.source == "stocktwits": social_msgs += 1 if r.raw_sentiment == 1.0: social_bull += 1 elif r.raw_sentiment == -1.0: social_bear += 1 agg.headline_count = len(rows) agg.primary_count = primary_count agg.first_headline_at = first_at.isoformat() if first_at else None agg.last_headline_at = last_at.isoformat() if last_at else None agg.category_counts = dict(cat_counter) agg.sentiment_mean = (sum(sentiments) / len(sentiments)) if sentiments else None agg.sentiment_recency_weighted = (weighted_num / weighted_den) if weighted_den > 0 else None directional = social_bull + social_bear agg.social = SocialStats( message_count=social_msgs, bull_count=social_bull, bear_count=social_bear, bull_bear_ratio=(social_bull / directional) if directional > 0 else None, ) agg.sources_present = sorted(sources) return agg