"""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 next_trading_day, trading_days_between from libs.labeler.reaction_date import compute_reaction_date logger = get_logger(__name__) LABEL_VERSION = "label-1.0.0" _LOOK_AHEAD_DAYS = 7 # fetch this many trading days of bars for 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 filing_time_bucket = getattr(event, "filing_time_bucket", "unknown") if not filing_time_bucket: 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 = next trading day after reaction_date 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: return EventLabel( event_id=event.event_id, entry_convention=entry_convention, reaction_date=reaction_date, entry_date=entry_date, label_status="unavailable", 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) < 5: 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) 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"), 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"), bars_to_mfe_3d=lbl_3d.get("bars_to_mfe"), bars_to_mae_3d=None, # symmetrically bars to MAE (min low) - optional 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, )