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.
609 lines
26 KiB
Python
609 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 if missing in DB feature_json.
|
|
# These can be None when the feature builder ran before reaction-day bars settled.
|
|
sym_bars = bars_by_symbol.get(sym, {})
|
|
rd = _parse_date(enriched.get("reaction_date"))
|
|
if rd and rd in sym_bars:
|
|
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 using config's scoring model for consistency with
|
|
# BacktestRunner. Only use config model when event_v1 features are
|
|
# present (parse_confidence_overall etc.), otherwise the model's hard
|
|
# gates reject events with incomplete features.
|
|
if enriched.get("parse_confidence_overall") is not None:
|
|
enriched["score"] = self._compute_score(enriched, config)
|
|
elif "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
|