"""Label Generator: compute forward-return labels for all valid events.""" from __future__ import annotations import argparse import asyncio import datetime as dt import uuid from sqlalchemy import select from libs.common.config import get_settings from libs.common.ids import new_job_run_id from libs.common.logging import bind_job_run_id, configure_logging, get_logger from libs.db.models import Event, EventLabel, JobRun, SymbolMaster from libs.db.session import get_session from libs.labeler.label_generator import LABEL_VERSION, generate_labels from libs.oracle_client import PriceService, make_oracle_client logger = get_logger(__name__) async def run_label_generator( run_id: str, entry_convention: str = "next_open_after_reaction_close", event_id_filter: str | None = None, start_date: str | None = None, end_date: str | None = None, ) -> dict[str, int]: stats = {"seen": 0, "labeled": 0, "skipped": 0, "errors": 0} async with get_session() as session, make_oracle_client() as oracle: price_svc = PriceService(oracle) job = JobRun( job_run_id=uuid.UUID(run_id), job_name="label_generator", source_name="oracle", run_date=dt.date.today(), status="running", ) session.add(job) await session.flush() # Query: events with status=valid (or specific event_id) stmt = select(Event, SymbolMaster).join( SymbolMaster, Event.symbol_id == SymbolMaster.symbol_id, isouter=True ).where(Event.status == "valid") if event_id_filter: stmt = stmt.where(Event.event_id == event_id_filter) if start_date: stmt = stmt.where(Event.event_date >= dt.date.fromisoformat(start_date)) if end_date: stmt = stmt.where(Event.event_date <= dt.date.fromisoformat(end_date)) result = await session.execute(stmt.order_by(Event.event_date, Event.event_id)) rows = result.all() stats["seen"] = len(rows) for event, symbol in rows: if symbol is None: logger.warning("label_no_symbol", event_id=event.event_id) stats["skipped"] += 1 continue # Skip if label already exists. Regenerate when: # - status == "pending" (market has since closed) # - status == "unavailable" AND entry_date is null/future (recover # from prior runs that 404'd on a future window — see # label_price_pending_future_window flow in label_generator). existing = await session.execute( select(EventLabel).where( EventLabel.event_id == event.event_id, EventLabel.entry_convention == entry_convention, EventLabel.label_version == LABEL_VERSION, ) ) existing_label = existing.scalar_one_or_none() if existing_label is not None: should_regenerate = existing_label.label_status == "pending" or ( existing_label.label_status == "unavailable" and ( existing_label.entry_date is None or existing_label.entry_date >= dt.date.today() ) ) if not should_regenerate: stats["skipped"] += 1 continue # Stale pending or recoverable unavailable — delete and regenerate await session.delete(existing_label) await session.flush() try: label = await generate_labels( session=session, event=event, price_svc=price_svc, ticker=symbol.ticker, entry_convention=entry_convention, ) session.add(label) await session.flush() stats["labeled"] += 1 logger.info( "label_created", event_id=event.event_id, label_status=label.label_status, ) except Exception as exc: logger.error("label_error", event_id=event.event_id, error=str(exc)) stats["errors"] += 1 job.status = "succeeded" if stats["errors"] == 0 else "partial" job.finished_at_utc = dt.datetime.now(tz=dt.UTC) job.records_seen = stats["seen"] job.records_written = stats["labeled"] job.error_count = stats["errors"] logger.info("label_generator_done", **stats) return stats def main() -> None: parser = argparse.ArgumentParser(description="Event Label Generator") parser.add_argument("--run-id", default=new_job_run_id()) parser.add_argument( "--entry-convention", default="next_open_after_reaction_close", choices=["next_open_after_reaction_close", "reaction_close"], help="Entry price convention", ) parser.add_argument("--event-id", default=None, help="Process a single event by ID") parser.add_argument("--start-date", default=None, metavar="YYYY-MM-DD") parser.add_argument("--end-date", default=None, metavar="YYYY-MM-DD") args = parser.parse_args() settings = get_settings() configure_logging(settings.log_level) bind_job_run_id(args.run_id) asyncio.run( run_label_generator( run_id=args.run_id, entry_convention=args.entry_convention, event_id_filter=args.event_id, start_date=args.start_date, end_date=args.end_date, ) ) if __name__ == "__main__": main()