""" Overlay pipeline orchestrator - coordinates data collection, feature building, and scoring for the full pipeline (batch) and on-demand (single symbol). """ import asyncio import logging from datetime import datetime, timedelta, timezone from typing import Dict, List, Optional from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select, desc from app.core.overlay_config import ONDEMAND_TIMEOUT_SECONDS, FEATURE_STALE_HOURS, TOP_50_SYMBOLS from app.models.overlay_feature import OverlayFeatureRecord, OverlayJobLog from app.services.overlay.yahoo_rss_adapter import YahooRSSAdapter from app.services.overlay.wikimedia_adapter import WikimediaAdapter from app.services.overlay.youtube_adapter import YouTubeAdapter from app.services.overlay.google_trends_adapter import GoogleTrendsAdapter from app.services.overlay.feature_builder import FeatureBuilder from app.services.overlay.overlay_scorer import OverlayScorer logger = logging.getLogger(__name__) class OverlayPipeline: """Orchestrate overlay data collection, feature building, and scoring.""" def __init__(self): self.rss = YahooRSSAdapter() self.wiki = WikimediaAdapter() self.youtube = YouTubeAdapter() self.trends = GoogleTrendsAdapter() self.builder = FeatureBuilder() self.scorer = OverlayScorer() # ------------------------------------------------------------------ # Job logging helpers # ------------------------------------------------------------------ async def _log_job( self, db: AsyncSession, job_type: str, status: str, started_at: datetime, records: int = 0, error: Optional[str] = None, ) -> None: log = OverlayJobLog( job_type=job_type, status=status, started_at=started_at, completed_at=datetime.now(timezone.utc) if status != "running" else None, records_processed=records, error_message=error, ) db.add(log) try: await db.commit() except Exception as e: logger.warning(f"Failed to persist job log: {e}") await db.rollback() # ------------------------------------------------------------------ # Data collection # ------------------------------------------------------------------ async def collect_all(self, db: AsyncSession) -> Dict: """Run all data collectors concurrently (source failures are isolated).""" started = datetime.now(timezone.utc) await self._log_job(db, "collect_all", "running", started) results = await asyncio.gather( self.rss.collect(db, symbols=TOP_50_SYMBOLS), self.wiki.collect(db), self.youtube.collect(db), self.trends.collect(db), return_exceptions=True, ) counts = { "yahoo_rss": results[0] if isinstance(results[0], int) else 0, "wikimedia": results[1] if isinstance(results[1], int) else 0, "youtube": results[2] if isinstance(results[2], int) else 0, "google_trends": results[3] if isinstance(results[3], int) else 0, } errors = [str(r) for r in results if isinstance(r, Exception)] error_str = "; ".join(errors) if errors else None final_status = "partial" if error_str else "completed" await self._log_job(db, "collect_all", final_status, started, sum(counts.values()), error_str) return counts # ------------------------------------------------------------------ # Feature building # ------------------------------------------------------------------ async def build_features_for_symbol( self, db: AsyncSession, symbol: str ) -> Optional[OverlayFeatureRecord]: """Build and persist a feature record for a single symbol.""" as_of = datetime.now(timezone.utc) features = await self.builder.build_all_features(db, symbol, as_of) scores = self.scorer.score(features) record = OverlayFeatureRecord( symbol=symbol.upper(), as_of_ts=as_of, feature_version="v1", # z-scores headline_burst_z=features.get("headline_burst_z"), youtube_influence_z=features.get("youtube_influence_z"), wiki_attention_z=features.get("wiki_attention_z"), theme_heat_z=features.get("theme_heat_z"), crowding_stress_z=features.get("crowding_stress_z"), # headline detail headline_count_6h=features.get("headline_count_6h", 0), headline_count_24h=features.get("headline_count_24h", 0), publisher_breadth_24h=features.get("publisher_breadth_24h", 0), # youtube detail youtube_mentions_24h=features.get("youtube_mentions_24h", 0), youtube_weighted_views_24h=features.get("youtube_weighted_views_24h", 0.0), # wiki detail wiki_views_1d=features.get("wiki_views_1d"), wiki_views_7d_avg=features.get("wiki_views_7d_avg"), # FINRA detail short_volume_ratio=features.get("short_volume_ratio"), short_volume_spike_zscore=features.get("short_volume_spike_zscore"), # scores **scores, ) db.add(record) await db.commit() await db.refresh(record) return record async def build_features_batch( self, db: AsyncSession, symbols: Optional[List[str]] = None ) -> int: """Build features for a batch of symbols (default: TOP_50_SYMBOLS).""" if symbols is None: symbols = TOP_50_SYMBOLS started = datetime.now(timezone.utc) await self._log_job(db, "feature_build", "running", started) built = 0 errors = [] for sym in symbols: try: await self.build_features_for_symbol(db, sym) built += 1 except Exception as e: logger.error(f"Feature build error for {sym}: {e}") errors.append(str(e)) error_str = "; ".join(errors[:3]) if errors else None final_status = "partial" if error_str else "completed" await self._log_job(db, "feature_build", final_status, started, built, error_str) return built # ------------------------------------------------------------------ # Query helpers # ------------------------------------------------------------------ async def get_latest_feature( self, db: AsyncSession, symbol: str ) -> Optional[OverlayFeatureRecord]: """Return the most recent feature record for a symbol.""" result = await db.execute( select(OverlayFeatureRecord) .where(OverlayFeatureRecord.symbol == symbol.upper()) .order_by(desc(OverlayFeatureRecord.as_of_ts)) .limit(1) ) return result.scalars().first() def is_stale(self, record: Optional[OverlayFeatureRecord]) -> bool: """Return True if the record is missing or older than FEATURE_STALE_HOURS.""" if record is None: return True # SQLite returns naive datetimes; ensure UTC-aware for comparison as_of = record.as_of_ts if as_of.tzinfo is None: as_of = as_of.replace(tzinfo=timezone.utc) age = datetime.now(timezone.utc) - as_of return age.total_seconds() > FEATURE_STALE_HOURS * 3600 # ------------------------------------------------------------------ # On-demand build # ------------------------------------------------------------------ async def get_or_build( self, db: AsyncSession, symbol: str ) -> Optional[OverlayFeatureRecord]: """ Return the latest feature record for *symbol*. If the record is stale/missing, trigger an on-demand pipeline run with a timeout guard. Returns the (possibly stale) record on timeout. """ record = await self.get_latest_feature(db, symbol) if record and not self.is_stale(record): return record logger.info(f"Overlay on-demand build triggered for {symbol}") try: record = await asyncio.wait_for( self.build_features_for_symbol(db, symbol), timeout=ONDEMAND_TIMEOUT_SECONDS, ) except asyncio.TimeoutError: logger.warning(f"Overlay on-demand build timed out for {symbol}") # Return potentially-stale record rather than None return record # ------------------------------------------------------------------ # Full pipeline # ------------------------------------------------------------------ async def run_full_pipeline(self, db: AsyncSession) -> Dict: """Run collect → build feature for all TOP_50 symbols.""" collect_counts = await self.collect_all(db) built = await self.build_features_batch(db) return { "collected": collect_counts, "features_built": built, }