You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

240 lines
8.0 KiB
Python

"""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,
)