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.

612 lines
26 KiB
Python

"""Detect new events from the pipeline DB and enrich them for paper trading.
Mirrors the enrichment pattern in SnapshotStore._async_load() but targets
a single execution_date rather than loading an entire parquet file.
"""
from __future__ import annotations
import asyncio
import datetime as dt
from typing import Any
from libs.backtest.domain import BacktestConfig
from libs.common.logging import get_logger
logger = get_logger(__name__)
class EventDetector:
"""Queries the pipeline DB for events executing on a given date and enriches them."""
def __init__(
self,
db_dsn: str,
oracle_url: str,
bars_cache: dict[str, dict[dt.date, dict]] | None = None,
) -> None:
self._db_dsn = db_dsn
self._oracle_url = oracle_url
self._bars_cache = bars_cache # pre-fetched bars from backtest_sim
self._db_unavailable: bool = False # circuit breaker: skip after first failure
self._screener_unavailable: bool = False # circuit breaker for screener API
self._company_cache: dict[str, dict[str, Any]] = {} # symbol -> {sector, market_cap}
self._screener_cache: dict[str, float | None] | None = None # cached screener results
@staticmethod
def _compute_score(row: dict[str, Any], config: BacktestConfig) -> float:
"""Compute score using the config's scoring_model — matches backtester."""
model = config.signal.scoring_model
if model == "return_max_long_v5":
from libs.backtest.scoring import compute_return_max_long_score_v5
return compute_return_max_long_score_v5(row)
elif model == "return_max_long_v7":
from libs.backtest.scoring import compute_return_max_long_score_v7
return compute_return_max_long_score_v7(row)
elif model == "return_max_long_v8":
from libs.backtest.scoring import compute_return_max_long_score_v8
return compute_return_max_long_score_v8(row)
elif model == "return_max_long_v9":
from libs.backtest.scoring import compute_return_max_long_score_v9
return compute_return_max_long_score_v9(row)
elif model == "return_max_long_v9g":
from libs.backtest.scoring import compute_return_max_long_score_v9g
return compute_return_max_long_score_v9g(row)
elif model == "return_max_long_v10":
from libs.backtest.scoring import compute_return_max_long_score_v10
return compute_return_max_long_score_v10(row)
elif model == "pead":
from libs.backtest.scoring import compute_pead_score
return compute_pead_score(row)
else:
from libs.backtest.scoring import compute_entry_score
return compute_entry_score(row)
async def get_candidates_for_date(
self,
execution_date: dt.date,
config: BacktestConfig,
convention: str | None = None,
) -> list[dict[str, Any]]:
"""Return enriched candidate rows for execution_date.
convention: 'reaction_close' for same-day close entries,
'next_open_after_reaction_close' for next-open entries.
None fetches both (legacy / debugging only).
Steps:
1. Query Event table for events with entry_date == execution_date
2. Fetch SymbolMaster tickers
3. Enrich with Alpaca price data (bars, avg_dollar_volume, ATR-14)
4. Enrich with Oracle sector data
5. Compute entry_price_est and score
"""
raw_rows = await self._fetch_events_for_date(execution_date, convention=convention)
if not raw_rows:
logger.debug("event_detector_no_events", date=execution_date.isoformat())
return []
unique_symbols = sorted({
str(r.get("symbol", "")).upper()
for r in raw_rows
if r.get("symbol")
})
# Fetch bars for the last 30 days for ADV + ATR computation
bar_start = execution_date - dt.timedelta(days=45)
bars_by_symbol, avg_dvol, atr_by_symbol = await self._fetch_enrichment_data(
unique_symbols, bar_start, execution_date
)
company_info = await self._fetch_company_info(unique_symbols)
screener_mcaps = await self._fetch_screener_market_caps(unique_symbols)
enriched_rows: list[dict[str, Any]] = []
for row in raw_rows:
sym = str(row.get("symbol", "")).upper()
if not sym:
continue
info = company_info.get(sym, {})
enriched = dict(row)
enriched["symbol"] = sym
# Use Oracle-computed avg_dvol; fall back to stored avg_dollar_volume_20d
# from feature_json when Oracle enrichment returned nothing.
oracle_adv = avg_dvol.get(sym)
enriched["avg_dollar_volume"] = (
oracle_adv if oracle_adv
else float(row.get("avg_dollar_volume_20d") or 0.0)
)
# Prefer atr_14 stored in DB feature_json (computed by feature builder with
# event_date+5d bars, giving stable post-event ATR). Fall back to Oracle-computed
# value only if not already present — avoids inflating ATR with the large
# earnings reaction-day range when bars end exactly at execution_date.
if not enriched.get("atr_14"):
enriched["atr_14"] = atr_by_symbol.get(sym)
# Enrich reaction_day_low / reaction_day_high from Oracle bars.
# These are NOT stored in DB feature_json but ARE used by compute_stop_price()
# to tighten stops when reaction_day_low > atr_stop (matches snapshot pipeline).
if not enriched.get("reaction_day_low") or not enriched.get("reaction_day_high"):
sym_bars = bars_by_symbol.get(sym, {})
reaction_date_raw = enriched.get("reaction_date")
if reaction_date_raw:
rd = _parse_date(reaction_date_raw)
if rd and rd in sym_bars:
if not enriched.get("reaction_day_low"):
enriched["reaction_day_low"] = sym_bars[rd].get("low")
if not enriched.get("reaction_day_high"):
enriched["reaction_day_high"] = sym_bars[rd].get("high")
# Compute market features from Oracle bars ONLY if missing in DB
# feature_json. DB values are authoritative because they were computed
# by the feature_builder at event time with the correct reaction_date
# and base price. Oracle bars can produce different values due to
# non-deterministic data or different date alignment.
_db_has_reaction = enriched.get("reaction_day_return") is not None
sym_bars = bars_by_symbol.get(sym, {})
rd = _parse_date(enriched.get("reaction_date"))
if rd and rd in sym_bars and not _db_has_reaction:
sorted_dates = sorted(sym_bars.keys())
try:
rd_idx = sorted_dates.index(rd)
except ValueError:
rd_idx = -1
if rd_idx > 0:
prior_dates = sorted_dates[max(0, rd_idx - 20):rd_idx]
prev_bar = sym_bars[sorted_dates[rd_idx - 1]]
reaction_bar = sym_bars[rd]
if not enriched.get("volume_ratio_20d") and prior_dates:
avg_vol = sum(sym_bars[d]["volume"] for d in prior_dates) / len(prior_dates)
if avg_vol > 0:
vr = reaction_bar["volume"] / avg_vol
enriched["volume_ratio_20d"] = vr
if not enriched.get("reaction_day_return") and prev_bar["close"]:
enriched["reaction_day_return"] = (
(reaction_bar["close"] - prev_bar["close"]) / prev_bar["close"]
)
if not enriched.get("close_location"):
rng = reaction_bar["high"] - reaction_bar["low"]
if rng > 0:
enriched["close_location"] = (
(reaction_bar["close"] - reaction_bar["low"]) / rng
)
if not enriched.get("gap_size") and prev_bar["close"]:
enriched["gap_size"] = (
(reaction_bar["open"] - prev_bar["close"]) / prev_bar["close"]
)
enriched["sector"] = info.get("sector", "UNKNOWN")
if "market_cap_proxy" not in enriched or enriched["market_cap_proxy"] is None:
# Prefer screener market_cap (same source as snapshot export pipeline)
market_cap = screener_mcaps.get(sym)
if market_cap is None:
market_cap = info.get("market_cap")
enriched["market_cap_proxy"] = market_cap
# entry_price_est: use reaction-day close price
if "entry_price_est" not in enriched or not enriched["entry_price_est"]:
event_close = enriched.get("event_close")
if event_close:
enriched["entry_price_est"] = float(event_close)
else:
# Fall back to latest available close from bars
sym_bars = bars_by_symbol.get(sym, {})
reaction_date_raw = enriched.get("reaction_date")
if reaction_date_raw:
rd = _parse_date(reaction_date_raw)
if rd and rd in sym_bars:
enriched["entry_price_est"] = sym_bars[rd].get("close", 0.0)
if not enriched.get("entry_price_est"):
# Final fallback: use the most recent available close
# (needed for same-day pending events where reaction_date bar not yet available)
sym_bars = bars_by_symbol.get(sym, {})
if sym_bars:
latest_d = max(sym_bars.keys())
enriched["entry_price_est"] = sym_bars[latest_d].get("close", 0.0)
if not enriched.get("entry_price_est"):
logger.debug(
"event_detector_skip_no_price",
symbol=sym,
event_id=enriched.get("event_id"),
)
continue
# Compute score if missing from DB.
# Use compute_entry_score (ranking-only, no hard gates).
# Trade filtering is handled by engine gates in select_candidates,
# not by the scoring function. Using v5/v8/v9 scoring here would
# reject too many candidates that engine gates should evaluate.
if "score" not in enriched or enriched.get("score") is None:
from libs.backtest.scoring import compute_entry_score
enriched["score"] = compute_entry_score(enriched)
enriched_rows.append(enriched)
logger.info(
"event_detector_candidates_ready",
date=execution_date.isoformat(),
count=len(enriched_rows),
)
return enriched_rows
# ------------------------------------------------------------------ #
# DB queries
# ------------------------------------------------------------------ #
async def _fetch_events_for_date(
self,
execution_date: dt.date,
convention: str | None = None,
) -> list[dict[str, Any]]:
"""Query EventLabel + FeatureSnapshot + Event + SymbolMaster for execution_date.
entry_date lives in EventLabel, not Event. Features live in FeatureSnapshot.feature_json.
convention: filter by entry_convention ('reaction_close' or 'next_open_after_reaction_close').
"""
if self._db_unavailable:
return []
try:
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from libs.db.models import Event, EventLabel, FeatureSnapshot, SymbolMaster
engine = create_async_engine(
self._db_dsn, echo=False,
connect_args={"timeout": 5},
)
async_session = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
_UTC = __import__("zoneinfo").ZoneInfo("UTC")
async with async_session() as session:
# Join: Event + EventLabel (entry_date) + FeatureSnapshot (features) + SymbolMaster
# Order by snapshot created_at ASC so newer snapshots override older ones when merging
stmt = (
select(Event, EventLabel, FeatureSnapshot, SymbolMaster)
.join(EventLabel, Event.event_id == EventLabel.event_id)
.join(FeatureSnapshot, Event.event_id == FeatureSnapshot.event_id)
.outerjoin(SymbolMaster, Event.symbol_id == SymbolMaster.symbol_id)
.where(EventLabel.entry_date == execution_date)
.where(EventLabel.label_status.in_(["ok", "truncated", "pending"]))
.order_by(FeatureSnapshot.created_at_utc.asc())
)
if convention is not None:
stmt = stmt.where(EventLabel.entry_convention == convention)
rows = (await session.execute(stmt)).all()
await engine.dispose()
# Group by event_id — merge ALL FeatureSnapshot.feature_json per event.
# The feature builder creates multiple snapshots per event with different feature types
# (market features, NLP features, fundamental features). We need all of them merged.
# Ordered ASC so newer keys override older ones.
event_meta: dict[str, tuple] = {} # event_id -> (event, label, sym)
event_features: dict[str, dict] = {} # event_id -> merged feature_json
for event, label, snapshot, sym in rows:
eid = event.event_id
if sym is None or not sym.ticker:
continue
if eid not in event_meta:
event_meta[eid] = (event, label, sym)
event_features[eid] = {}
# Merge: newer snapshot keys override older ones (ASC order)
event_features[eid].update(snapshot.feature_json or {})
result: list[dict[str, Any]] = []
for eid, (event, label, sym) in event_meta.items():
ts = event.filed_at_utc
if ts is None and event.event_date is not None:
ts = dt.datetime.combine(
event.event_date, dt.time(21, 0), tzinfo=_UTC
)
row: dict[str, Any] = {
"event_id": eid,
"symbol": sym.ticker,
"issuer_id": event.issuer_id,
"event_type": event.event_type or "",
"event_direction": event.event_direction or "",
"event_date": event.event_date,
"event_timestamp": ts,
"reaction_date": label.reaction_date,
"entry_date": label.entry_date,
"execution_date": execution_date,
"label_status": label.label_status,
# Merged features from ALL FeatureSnapshots for this event
**event_features[eid],
}
result.append(row)
logger.debug(
"event_detector_db_fetched",
date=execution_date.isoformat(),
count=len(result),
)
return result
except Exception as exc:
if not self._db_unavailable:
err_msg = f"{type(exc).__name__}: {exc}" if str(exc) else type(exc).__name__
logger.warning("event_detector_db_fetch_failed", error=err_msg)
self._db_unavailable = True
return []
# ------------------------------------------------------------------ #
# Enrichment helpers
# ------------------------------------------------------------------ #
async def _fetch_enrichment_data(
self,
symbols: list[str],
start: dt.date,
end: dt.date,
concurrency: int = 16,
) -> tuple[
dict[str, dict[dt.date, dict[str, Any]]], # bars_by_symbol
dict[str, float], # avg_dvol
dict[str, float | None], # atr_14
]:
if not symbols:
return {}, {}, {}
# Fast path: use pre-fetched bars_cache (backtest mode)
if self._bars_cache is not None:
return self._enrichment_from_cache(symbols, start, end)
try:
from libs.oracle_client import OracleClient, PriceService
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]] = {}
avg_dvol: dict[str, float] = {}
atr_14_map: dict[str, float | None] = {}
semaphore = asyncio.Semaphore(concurrency)
async with OracleClient(base_url=self._oracle_url) as client:
svc = PriceService(client)
async def _fetch_one(sym: str) -> None:
async with semaphore:
try:
resp = await svc.get_daily_bars(
sym,
start=start.isoformat(),
end=end.isoformat(),
)
date_map: dict[dt.date, dict[str, Any]] = {}
dollar_vols: list[float] = []
closes: list[float] = []
for bar in resp.bars:
d = dt.date.fromisoformat(bar.date)
b = {
"date": d,
"open": bar.open,
"high": bar.high,
"low": bar.low,
"close": bar.close,
"volume": bar.volume,
}
date_map[d] = b
dollar_vols.append(bar.close * bar.volume)
closes.append(bar.close)
bars_by_symbol[sym] = date_map
last_20 = dollar_vols[-20:]
avg_dvol[sym] = sum(last_20) / len(last_20) if last_20 else 0.0
atr_14_map[sym] = _compute_atr14_from_dicts(list(date_map.values()))
except Exception as sym_exc:
err_msg = f"{type(sym_exc).__name__}: {sym_exc}" if str(sym_exc) else type(sym_exc).__name__
# "Not found" is expected for delisted/renamed symbols — debug only
if "not found" in str(sym_exc).lower():
logger.debug("event_detector_price_fetch_skipped", symbol=sym, error=err_msg)
else:
logger.warning("event_detector_price_fetch_failed", symbol=sym, error=err_msg)
bars_by_symbol[sym] = {}
avg_dvol[sym] = 0.0
atr_14_map[sym] = None
await asyncio.gather(*(_fetch_one(sym) for sym in symbols))
return bars_by_symbol, avg_dvol, atr_14_map
except Exception as exc:
err_msg = f"{type(exc).__name__}: {exc}" if str(exc) else type(exc).__name__
logger.warning("event_detector_oracle_failed", error=err_msg)
return {}, {}, {}
def _enrichment_from_cache(
self,
symbols: list[str],
start: dt.date,
end: dt.date,
) -> tuple[
dict[str, dict[dt.date, dict[str, Any]]],
dict[str, float],
dict[str, float | None],
]:
"""Compute enrichment data from pre-fetched bars_cache (no Oracle calls)."""
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]] = {}
avg_dvol: dict[str, float] = {}
atr_14_map: dict[str, float | None] = {}
for sym in symbols:
all_bars = self._bars_cache.get(sym, {})
# Filter to requested date range
date_map = {d: b for d, b in all_bars.items() if start <= d <= end}
bars_by_symbol[sym] = date_map
if not date_map:
avg_dvol[sym] = 0.0
atr_14_map[sym] = None
continue
sorted_bars = [date_map[d] for d in sorted(date_map.keys())]
dollar_vols = [b["close"] * b["volume"] for b in sorted_bars]
last_20 = dollar_vols[-20:]
avg_dvol[sym] = sum(last_20) / len(last_20) if last_20 else 0.0
atr_14_map[sym] = _compute_atr14_from_dicts(sorted_bars)
return bars_by_symbol, avg_dvol, atr_14_map
async def _fetch_screener_market_caps(
self,
symbols: list[str],
) -> dict[str, float | None]:
"""Fetch market_cap for symbols via Oracle screener (batch).
Mirrors snapshot_export._resolve_universe_profile() which uses
ScreenerService to get market_cap_proxy for all universe stocks.
"""
if not symbols or self._screener_unavailable:
return {}
# Return from cache if available (screener data is static across days)
if self._screener_cache is not None:
symbol_set = {s.upper() for s in symbols}
return {s: self._screener_cache.get(s) for s in symbol_set}
symbol_set = {s.upper() for s in symbols}
result: dict[str, float | None] = {s: None for s in symbol_set}
try:
from libs.oracle_client import OracleClient, ScreenerService
async with OracleClient(base_url=self._oracle_url) as client:
svc = ScreenerService(client)
# market_cap_min=500M matches snapshot_export pattern: reduces result set
# from ~10k to ~3k stocks, avoiding HTTP 500 on large responses.
stocks = await svc.search_all_stocks(
market_cap_min=500_000_000,
exchange="NYSE,NASDAQ,AMEX",
exclude_types="ETF,FUND,ADR,SPAC",
)
# Cache ALL screener results for future calls
self._screener_cache = {}
for stock in stocks:
sym = (stock.symbol or "").upper()
self._screener_cache[sym] = stock.market_cap
if sym in symbol_set:
result[sym] = stock.market_cap
except Exception as exc:
err_msg = f"{type(exc).__name__}: {exc}" if str(exc) else type(exc).__name__
logger.warning("event_detector_screener_mcap_failed", error=err_msg)
self._screener_unavailable = True
return result
async def _fetch_company_info(
self,
symbols: list[str],
concurrency: int = 16,
) -> dict[str, dict[str, Any]]:
"""Fetch sector and market_cap for each symbol from Oracle (cached)."""
if not symbols:
return {}
# Only fetch symbols not yet in cache
uncached = [s for s in symbols if s not in self._company_cache]
if uncached:
semaphore = asyncio.Semaphore(concurrency)
try:
from libs.oracle_client import CompanyService, OracleClient
async with OracleClient(base_url=self._oracle_url) as client:
company_svc = CompanyService(client)
async def _fetch_one(sym: str) -> tuple[str, dict[str, Any]]:
async with semaphore:
try:
info = await company_svc.get_company(sym)
return sym, {
"sector": info.sector or "UNKNOWN",
"market_cap": info.market_cap,
}
except Exception:
return sym, {"sector": "UNKNOWN", "market_cap": None}
fetched = await asyncio.gather(*(_fetch_one(sym) for sym in uncached))
for sym, info in fetched:
self._company_cache[sym] = info
except Exception as exc:
err_msg = f"{type(exc).__name__}: {exc}" if str(exc) else type(exc).__name__
logger.warning("event_detector_company_fetch_failed", error=err_msg)
result: dict[str, dict[str, Any]] = {}
for sym in symbols:
result[sym] = self._company_cache.get(sym, {"sector": "UNKNOWN", "market_cap": None})
return result
# ------------------------------------------------------------------ #
# Helpers
# ------------------------------------------------------------------ #
def _compute_atr14(bars: list[Any]) -> float | None:
"""Compute 14-day ATR from a list of bar objects (API response)."""
if len(bars) < 2:
return None
true_ranges: list[float] = []
for i in range(1, len(bars)):
prev_close = float(bars[i - 1].close)
high = float(bars[i].high)
low = float(bars[i].low)
tr = max(
high - low,
abs(high - prev_close),
abs(low - prev_close),
)
true_ranges.append(tr)
if not true_ranges:
return None
last_14 = true_ranges[-14:] if len(true_ranges) >= 14 else true_ranges
return sum(last_14) / len(last_14)
def _compute_atr14_from_dicts(bars: list[dict]) -> float | None:
"""Compute 14-day ATR from a list of bar dicts (cache format)."""
if len(bars) < 2:
return None
true_ranges: list[float] = []
for i in range(1, len(bars)):
prev_close = float(bars[i - 1]["close"])
high = float(bars[i]["high"])
low = float(bars[i]["low"])
tr = max(
high - low,
abs(high - prev_close),
abs(low - prev_close),
)
true_ranges.append(tr)
if not true_ranges:
return None
last_14 = true_ranges[-14:] if len(true_ranges) >= 14 else true_ranges
return sum(last_14) / len(last_14)
def _parse_date(raw: Any) -> dt.date | None:
if isinstance(raw, dt.datetime):
return raw.date()
if isinstance(raw, dt.date):
return raw
if isinstance(raw, str):
try:
return dt.date.fromisoformat(raw[:10])
except ValueError:
return None
return None