"""Generate forward-return labels for parsed events using price bars.""" from __future__ import annotations import datetime as dt from decimal import Decimal from typing import Any from sqlalchemy.ext.asyncio import AsyncSession from libs.common.logging import get_logger from libs.common.time_utils import filing_time_bucket as classify_time_bucket, next_trading_day, trading_days_between from libs.labeler.reaction_date import compute_reaction_date logger = get_logger(__name__) LABEL_VERSION = "label-2.0.0" _LOOK_AHEAD_DAYS = 25 # fetch enough trading days for 20D label computation _R_FACTOR = 0.01 # 1R = 1% move (used for hit_pos/neg_1r labels) def _safe_decimal(v: Any) -> Decimal | None: try: return Decimal(str(v)) if v is not None else None except Exception: return None def _pct_return(entry: Decimal, exit_price: Decimal) -> Decimal: if entry == 0: return Decimal("0") return (exit_price - entry) / entry def _compute_labels_from_bars( entry_price: Decimal, bars: list[dict[str, Any]], # sorted ascending by date n_days: int, ) -> dict[str, Any]: """Compute forward-return labels over `n_days` trading days of bars. Args: entry_price: Entry price (open of entry date). bars: List of OHLCV dicts with keys: date, open, high, low, close. n_days: Look-ahead horizon (3 or 5). Returns: Dict of computed label fields for the given horizon. """ window = bars[:n_days] if not window: return {} closes = [_safe_decimal(b.get("close")) for b in window] highs = [_safe_decimal(b.get("high")) for b in window] lows = [_safe_decimal(b.get("low")) for b in window] # Forward close return at day n last_close = closes[-1] fwd_return = _pct_return(entry_price, last_close) if last_close else None # MFE (max favorable excursion): max high vs entry price valid_highs = [h for h in highs if h is not None] mfe = (max(valid_highs) - entry_price) / entry_price if valid_highs else None # MAE (max adverse excursion): min low vs entry price valid_lows = [lo for lo in lows if lo is not None] mae = (min(valid_lows) - entry_price) / entry_price if valid_lows else None # Hit +1R within n days threshold_pos = entry_price * (1 + Decimal(str(_R_FACTOR))) hit_pos = any(h is not None and h >= threshold_pos for h in highs) # Hit -1R within n days threshold_neg = entry_price * (1 - Decimal(str(_R_FACTOR))) hit_neg = any(lo is not None and lo <= threshold_neg for lo in lows) # Close up after n days close_up = bool(last_close is not None and last_close > entry_price) # Bars to MFE (index of max high) bars_to_mfe: int | None = None if valid_highs and mfe is not None: max_high = max(valid_highs) for i, h in enumerate(highs): if h == max_high: bars_to_mfe = i + 1 break # Days to peak close valid_close_idx = [(i, c) for i, c in enumerate(closes) if c is not None] days_to_peak_close: int | None = None if valid_close_idx: peak_close_idx = max(valid_close_idx, key=lambda x: x[1])[0] days_to_peak_close = peak_close_idx + 1 return { "fwd_return": fwd_return, "mfe": mfe, "mae": mae, "hit_pos_1r": hit_pos, "hit_neg_1r": hit_neg, "close_up": close_up, "bars_to_mfe": bars_to_mfe, "days_to_peak_close": days_to_peak_close, } async def generate_labels( session: AsyncSession, event: Any, # Event ORM model instance price_svc: Any, # PriceService ticker: str, entry_convention: str = "next_open_after_reaction_close", ) -> Any: """Generate EventLabel for a single event. Args: session: Async DB session. event: Event ORM instance (needs event_date, filed_at_utc, symbol_id). price_svc: PriceService instance for fetching bars. ticker: Trading ticker symbol. entry_convention: How to determine entry price. Returns: EventLabel ORM instance (not yet added to session). """ from libs.db.models import EventLabel # 1. Compute reaction_date if event.filed_at_utc is not None: filing_time_bucket = classify_time_bucket(event.filed_at_utc) else: filing_time_bucket = "unknown" event_date: dt.date = event.event_date reaction_date = compute_reaction_date(event_date, filing_time_bucket) # 2. Compute entry_date # reaction_close: enter at close of reaction_date itself # next_open_after_reaction_close: enter at open of next trading day if entry_convention == "reaction_close": entry_date = reaction_date else: entry_date = next_trading_day(reaction_date) # 3. Fetch price bars (entry_date + _LOOK_AHEAD_DAYS trading days) fetch_start = entry_date trading_days = trading_days_between(entry_date, entry_date + dt.timedelta(days=20)) fetch_end = trading_days[_LOOK_AHEAD_DAYS] if len(trading_days) > _LOOK_AHEAD_DAYS else trading_days[-1] try: price_resp = await price_svc.get_daily_bars( ticker=ticker, start=fetch_start.isoformat(), end=fetch_end.isoformat(), ) raw_bars = [b.model_dump() for b in price_resp.bars] except Exception as exc: logger.warning("label_price_unavailable", ticker=ticker, error=str(exc)) return EventLabel( event_id=event.event_id, entry_convention=entry_convention, reaction_date=reaction_date, entry_date=None, label_status="unavailable", invalid_event_for_labeling=False, label_version=LABEL_VERSION, ) # Filter bars from entry_date onwards, sorted ascending bars = sorted( [b for b in raw_bars if b.get("date") and b["date"] >= entry_date.isoformat()], key=lambda b: b["date"], ) if not bars: # For reaction_close, the bar may not exist yet if market hasn't closed. # Return "pending" so the pipeline can regenerate after market close. if entry_convention == "reaction_close" and entry_date >= dt.date.today(): label_status_no_bar = "pending" else: label_status_no_bar = "unavailable" return EventLabel( event_id=event.event_id, entry_convention=entry_convention, reaction_date=reaction_date, entry_date=entry_date, label_status=label_status_no_bar, invalid_event_for_labeling=False, label_version=LABEL_VERSION, ) # 4. Determine entry_price first_bar = bars[0] if entry_convention == "next_open_after_reaction_close": entry_price = _safe_decimal(first_bar.get("open")) else: entry_price = _safe_decimal(first_bar.get("close")) if entry_price is None or entry_price == 0: return EventLabel( event_id=event.event_id, entry_convention=entry_convention, reaction_date=reaction_date, entry_date=entry_date, entry_price=None, label_status="unavailable", invalid_event_for_labeling=True, label_version=LABEL_VERSION, ) # 5. Forward bars (exclude entry bar itself for 1D/3D/5D) forward_bars = bars[1:] # Day 1+ after entry label_status = "ok" if len(forward_bars) < 20: label_status = "truncated" # 1D return fwd_1d = _pct_return(entry_price, _safe_decimal(bars[1]["close"])) if len(bars) > 1 else None # 3D labels lbl_3d = _compute_labels_from_bars(entry_price, forward_bars, 3) # 5D labels lbl_5d = _compute_labels_from_bars(entry_price, forward_bars, 5) # 10D labels lbl_10d = _compute_labels_from_bars(entry_price, forward_bars, 10) # 20D labels lbl_20d = _compute_labels_from_bars(entry_price, forward_bars, 20) return EventLabel( event_id=event.event_id, entry_convention=entry_convention, reaction_date=reaction_date, entry_date=entry_date, entry_price=entry_price, fwd_return_1d=fwd_1d, fwd_return_3d=lbl_3d.get("fwd_return"), fwd_return_5d=lbl_5d.get("fwd_return"), fwd_return_10d=lbl_10d.get("fwd_return"), fwd_return_20d=lbl_20d.get("fwd_return"), hit_pos_1r_within_3d=lbl_3d.get("hit_pos_1r"), hit_neg_1r_within_3d=lbl_3d.get("hit_neg_1r"), close_up_after_3d=lbl_3d.get("close_up"), close_up_after_5d=lbl_5d.get("close_up"), mfe_3d=lbl_3d.get("mfe"), mae_3d=lbl_3d.get("mae"), mfe_5d=lbl_5d.get("mfe"), mae_5d=lbl_5d.get("mae"), mfe_10d=lbl_10d.get("mfe"), mae_10d=lbl_10d.get("mae"), mfe_20d=lbl_20d.get("mfe"), mae_20d=lbl_20d.get("mae"), bars_to_mfe_3d=lbl_3d.get("bars_to_mfe"), bars_to_mae_3d=None, days_to_peak_close_5d=lbl_5d.get("days_to_peak_close"), label_status=label_status, invalid_event_for_labeling=False, label_version=LABEL_VERSION, )