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.

111 lines
3.4 KiB
Python

"""Market-side feature calculations from price bar data."""
from __future__ import annotations
from typing import Any
from libs.oracle_client.models import PriceBar
def reaction_day_return(bars: list[PriceBar], event_date: str) -> float | None:
"""(close - prev_close) / prev_close on event date."""
dated = {b.date: b for b in bars}
if event_date not in dated:
return None
event_bar = dated[event_date]
# Find previous bar
sorted_dates = sorted(dated.keys())
idx = sorted_dates.index(event_date)
if idx == 0:
return None
prev_bar = dated[sorted_dates[idx - 1]]
if prev_bar.close == 0:
return None
return (event_bar.close - prev_bar.close) / prev_bar.close
def volume_ratio_20d(bars: list[PriceBar], event_date: str) -> float | None:
"""Event day volume / 20-day average volume before event."""
dated = {b.date: b for b in bars}
sorted_dates = sorted(dated.keys())
if event_date not in dated:
return None
idx = sorted_dates.index(event_date)
if idx < 1:
return None
prior = sorted_dates[max(0, idx - 20) : idx]
if not prior:
return None
avg_vol = sum(dated[d].volume for d in prior) / len(prior)
if avg_vol == 0:
return None
return dated[event_date].volume / avg_vol
def close_location(bar: PriceBar) -> float | None:
"""(close - low) / (high - low): 0=closed at low, 1=at high."""
rng = bar.high - bar.low
if rng == 0:
return None
return (bar.close - bar.low) / rng
def gap_size(bars: list[PriceBar], event_date: str) -> float | None:
"""(open_today - close_yesterday) / close_yesterday."""
dated = {b.date: b for b in bars}
sorted_dates = sorted(dated.keys())
if event_date not in dated:
return None
idx = sorted_dates.index(event_date)
if idx == 0:
return None
today = dated[event_date]
yesterday = dated[sorted_dates[idx - 1]]
if yesterday.close == 0:
return None
return (today.open - yesterday.close) / yesterday.close
def atr_14(bars: list[PriceBar]) -> float | None:
"""14-period Average True Range."""
if len(bars) < 2:
return None
sorted_bars = sorted(bars, key=lambda b: b.date)
true_ranges: list[float] = []
for i in range(1, len(sorted_bars)):
curr = sorted_bars[i]
prev = sorted_bars[i - 1]
tr = max(
curr.high - curr.low,
abs(curr.high - prev.close),
abs(curr.low - prev.close),
)
true_ranges.append(tr)
if len(true_ranges) < 14:
return sum(true_ranges) / len(true_ranges) if true_ranges else None
return sum(true_ranges[-14:]) / 14
def compute_market_features(
bars: list[PriceBar], event_date: str
) -> dict[str, Any]:
"""Compute all market features for an event date."""
dated = {b.date: b for b in bars}
event_bar = dated.get(event_date)
features: dict[str, Any] = {
"reaction_day_return": reaction_day_return(bars, event_date),
"volume_ratio_20d": volume_ratio_20d(bars, event_date),
"gap_size": gap_size(bars, event_date),
"atr_14": atr_14(bars),
}
if event_bar:
features["close_location"] = close_location(event_bar)
features["event_close"] = event_bar.close
features["event_volume"] = event_bar.volume
else:
features["close_location"] = None
features["event_close"] = None
features["event_volume"] = None
return features