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.

2362 lines
100 KiB
Python

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""Point-in-time snapshot data loader for the backtester.
The public interface is:
store = SnapshotStore.load(snapshot_dir, split_name, oracle_url, db_dsn)
rows = store.get_candidates_for_date(date)
bar = store.get_bar(symbol, date)
macro = store.get_macro_for_date(date)
For tests, inject data directly:
store = SnapshotStore(candidates_by_exec_date=..., bars_by_symbol_date=..., macro_by_date=...)
"""
from __future__ import annotations
import asyncio
import datetime as dt
import hashlib
import json
import math
import os
from functools import partial
import pickle
from pathlib import Path
import time
from typing import Any
from zoneinfo import ZoneInfo
import pyarrow as pa
import pyarrow.parquet as pq
from libs.backtest.proxies import peer_candidates_for_symbol, sector_etf_for_sector
from libs.features.market_features import compute_market_features
from libs.common.config import get_settings
from libs.common.logging import get_logger
logger = get_logger(__name__)
_UTC = ZoneInfo("UTC")
_PRICE_DERIVED_FEATURE_COLUMNS = (
"pre_event_volatility_20d",
"pre_event_rsi_14",
"pre_event_bb_position",
"pre_event_obv_slope_20d",
"pre_event_hurst_60d",
"pre_event_entropy_60d",
"pre_event_ou_theta_60d",
"pre_event_gravitational_pull",
"pre_event_market_temperature",
)
_MACRO_BACKFILL_FEATURE_COLUMNS = (
"macro_vix",
"macro_hy_spread",
"macro_t10y2y",
)
_MACRO_TRADE_BAR_PREFIXES: tuple[tuple[str, str], ...] = (
("SPY", "spy"),
("QQQ", "qqq"),
("SH", "sh"),
("PSQ", "psq"),
("GLD", "gld"),
("SPYM", "spym"),
("QUAL", "qual"),
("QQQM", "qqqm"),
("TQQQ", "tqqq"),
("SHY", "shy"),
("USFR", "usfr"),
("BIL", "bil"),
("VGSH", "vgsh"),
("IEI", "iei"),
("IEF", "ief"),
("TIP", "tip"),
("XLK", "xlk"),
("SMH", "smh"),
("IWM", "iwm"),
("XLI", "xli"),
("XLE", "xle"),
("XLF", "xlf"),
("DBC", "dbc"),
("SGOV", "sgov"),
("JEPQ", "jepq"),
("BUFB", "bufb"),
("MERIX", "merix"),
)
_PRICE_FEATURE_WARMUP_DAYS = 400 # bumped from 120 to cover xsmom 273-trading-day lookback (12+1 momentum). Affects bars-fetch only; other engines unaffected.
_RUNTIME_CACHE_VERSION = 13
_RUNTIME_CACHE_WAIT_TIMEOUT_SECONDS = 1800.0
_RUNTIME_CACHE_WAIT_LOG_INTERVAL_SECONDS = 5.0
_EVENT_METADATA_FETCH_CHUNK_SIZE = 10_000
class SnapshotStore:
"""In-memory cache of point-in-time backtest data."""
def __init__(
self,
candidates_by_exec_date: dict[dt.date, list[dict[str, Any]]],
bars_by_symbol_date: dict[str, dict[dt.date, dict[str, Any]]],
macro_by_date: dict[dt.date, dict[str, Any]] | None = None,
) -> None:
self._candidates = {
date: list(rows)
for date, rows in candidates_by_exec_date.items()
}
self._candidates_by_reaction_date = self._build_reaction_index(self._candidates)
self._bars = bars_by_symbol_date
self._macro = macro_by_date or {}
self._price_bar_cache: dict[str, list[Any]] = {}
self._market_feature_cache: dict[tuple[str, dt.date], dict[str, Any]] = {}
self.snapshot_dir: Path | None = None
# ------------------------------------------------------------------
# Public query interface
# ------------------------------------------------------------------
def get_candidates_for_date(self, date: dt.date) -> list[dict[str, Any]]:
"""Return candidates where execution_date == date. No look-ahead."""
return list(self._candidates.get(date, []))
def get_candidates_for_reaction_date(self, date: dt.date) -> list[dict[str, Any]]:
"""Return candidates where reaction_date == date. No look-ahead."""
return list(self._candidates_by_reaction_date.get(date, []))
def get_bar(self, symbol: str, date: dt.date) -> dict[str, Any] | None:
"""Return OHLCV bar for symbol on date, or None if missing."""
sym_bars = self._bars.get(symbol)
if sym_bars is None:
return None
return sym_bars.get(date)
def get_latest_bar_on_or_before(
self,
symbol: str,
date: dt.date,
) -> tuple[dt.date, dict[str, Any]] | None:
"""Return the latest available bar for ``symbol`` on or before ``date``."""
sym_bars = self._bars.get(symbol)
if not sym_bars:
return None
eligible_dates = [bar_date for bar_date in sym_bars if bar_date <= date]
if not eligible_dates:
return None
latest_date = max(eligible_dates)
return latest_date, sym_bars[latest_date]
def get_macro_for_date(self, date: dt.date) -> dict[str, Any]:
"""Return macro observations for date (empty dict if none)."""
return dict(self._macro.get(date, {}))
def get_market_features(self, symbol: str, event_date: dt.date) -> dict[str, Any]:
"""Return market-side features for a stored symbol as of event_date."""
cache_key = (str(symbol).upper(), event_date)
if cache_key in self._market_feature_cache:
return dict(self._market_feature_cache[cache_key])
price_bars = self._get_price_bars_for_symbol(
str(symbol).upper(),
bars_by_symbol=self._bars,
price_bar_cache=self._price_bar_cache,
)
if not price_bars:
self._market_feature_cache[cache_key] = {}
return {}
features = compute_market_features(price_bars, event_date.isoformat())
self._market_feature_cache[cache_key] = dict(features)
return dict(features)
def all_execution_dates(self) -> list[dt.date]:
"""Sorted list of dates that have at least one candidate."""
return sorted(self._candidates.keys())
def all_reaction_dates(self) -> list[dt.date]:
"""Sorted list of dates that have at least one reaction-date candidate."""
return sorted(self._candidates_by_reaction_date.keys())
def all_trading_days(self, include_reaction_dates: bool = False) -> list[dt.date]:
"""All NYSE trading days from first execution date to last available date.
Extends beyond the last execution date using macro data (SPY/QQQ prices)
so parking can run on days with no events. Warm-up dates before the first
execution date are excluded.
"""
from libs.backtest.calendar import get_trading_days
dates = set(self.all_execution_dates())
if include_reaction_dates:
dates.update(self.all_reaction_dates())
# Extend END of range with macro dates (for parking after last event)
if self._macro and dates:
min_exec = min(dates)
macro_after = {d for d in self._macro if d >= min_exec}
dates.update(macro_after)
elif self._macro and not dates:
# No events at all (parking_only) — use full macro range
dates.update(self._macro.keys())
if not dates:
return []
ordered = sorted(dates)
last_day = ordered[-1]
# Cap at the last day where parking symbols actually have close-price data.
# SPY/QQQ are fetched as core and always present; QQQM/TQQQ/SGOV are fetched
# as extras and can lag by minutes after market close. Using max(macro.keys())
# would pick a day where SPY data exists but parking prices are missing.
if self._macro:
_parking_close_keys = ("qqqm_close", "tqqq_close", "sgov_close")
_per_symbol_last: list[dt.date] = []
for _k in _parking_close_keys:
_sym_dates = [d for d in self._macro if self._macro[d].get(_k)]
if _sym_dates:
_per_symbol_last.append(max(_sym_dates))
if _per_symbol_last:
last_day = min(last_day, min(_per_symbol_last))
else:
last_day = min(last_day, max(self._macro.keys()))
return get_trading_days(ordered[0], last_day)
def slice_by_date_range(
self,
start_date: dt.date,
end_date: dt.date,
*,
require_reaction_date_in_range: bool = False,
clamp_reaction_index_to_window: bool = True,
) -> "SnapshotStore":
"""Return a copy restricted to candidates within a date window.
Candidates are always filtered by execution date. When
``require_reaction_date_in_range`` is True, rows with a reaction date
outside the requested window are excluded as well. By default rows are
kept if they execute in-range, but the reaction-date index is clamped
to the requested window so reaction-close simulations do not expand the
fold backwards. Bars are kept intact, while macro observations are
trimmed to the same date range.
"""
filtered_candidates: dict[dt.date, list[dict[str, Any]]] = {}
for exec_date, rows in self._candidates.items():
if exec_date < start_date or exec_date > end_date:
continue
kept_rows: list[dict[str, Any]] = []
for row in rows:
if require_reaction_date_in_range:
reaction_date = self._normalize_date(row.get("reaction_date"))
if reaction_date is not None and (
reaction_date < start_date or reaction_date > end_date
):
continue
kept_rows.append(dict(row))
if kept_rows:
filtered_candidates[exec_date] = kept_rows
# Keep macro data beyond end_date so parking can run after last event.
# Only trim dates before start_date (warm-up excluded).
filtered_macro = {
date: dict(values)
for date, values in self._macro.items()
if date >= start_date
}
sliced = SnapshotStore(
candidates_by_exec_date=filtered_candidates,
bars_by_symbol_date=self._bars,
macro_by_date=filtered_macro,
)
if clamp_reaction_index_to_window:
sliced._candidates_by_reaction_date = {
date: list(rows)
for date, rows in sliced._candidates_by_reaction_date.items()
if start_date <= date <= end_date
}
sliced.snapshot_dir = self.snapshot_dir
return sliced
# ------------------------------------------------------------------
# Factory: load from Parquet + DB + Oracle
# ------------------------------------------------------------------
@classmethod
def load(
cls,
snapshot_dir: str | Path,
split_name: str,
oracle_url: str,
db_dsn: str,
scoring_fn: Any | None = None,
) -> "SnapshotStore":
"""Synchronous factory. Internally uses asyncio.run() to prefetch data.
scoring_fn: optional callable(row_dict) -> float to override default scoring.
Raises RuntimeError if called from within a running event loop.
"""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop is not None and loop.is_running():
raise RuntimeError(
"SnapshotStore.load() cannot be called from a running event loop. "
"Use SnapshotStore._async_load() directly in async contexts."
)
snapshot_path = Path(snapshot_dir)
store = cls._load_with_runtime_cache(
snapshot_path=snapshot_path,
split_names=[split_name],
scoring_fn=scoring_fn,
builder=lambda: asyncio.run(
cls._async_load(snapshot_path, split_name, oracle_url, db_dsn, scoring_fn)
),
)
store.snapshot_dir = snapshot_path
return store
@classmethod
def load_merged(
cls,
snapshot_dir: str | Path,
split_names: list[str],
oracle_url: str,
db_dsn: str,
scoring_fn: Any | None = None,
) -> "SnapshotStore":
"""Load multiple splits in one pass to avoid duplicate DB/Oracle fetches."""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop is not None and loop.is_running():
raise RuntimeError(
"SnapshotStore.load_merged() cannot be called from a running event loop. "
"Use SnapshotStore._async_load_merged() directly in async contexts."
)
snapshot_path = Path(snapshot_dir)
store = cls._load_with_runtime_cache(
snapshot_path=snapshot_path,
split_names=split_names,
scoring_fn=scoring_fn,
builder=lambda: asyncio.run(
cls._async_load_merged(snapshot_path, split_names, oracle_url, db_dsn, scoring_fn)
),
)
store.snapshot_dir = snapshot_path
return store
@classmethod
def materialize_snapshot_dir(
cls,
snapshot_dir: str | Path,
oracle_url: str,
db_dsn: str,
*,
split_names: list[str] | None = None,
output_dir: str | Path | None = None,
) -> list[str]:
"""Persist runtime-backfilled feature columns into snapshot parquet files."""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop is not None and loop.is_running():
raise RuntimeError(
"SnapshotStore.materialize_snapshot_dir() cannot be called from a running event loop. "
"Use SnapshotStore._async_materialize_snapshot_dir() directly in async contexts."
)
return asyncio.run(
cls._async_materialize_snapshot_dir(
Path(snapshot_dir),
oracle_url=oracle_url,
db_dsn=db_dsn,
split_names=split_names,
output_dir=Path(output_dir) if output_dir is not None else None,
)
)
@classmethod
async def _async_load(
cls,
snapshot_dir: Path,
split_name: str,
oracle_url: str,
db_dsn: str,
scoring_fn: Any | None = None,
) -> dict[str, Any]:
"""Async data loading pipeline. Returns kwargs dict for __init__."""
parquet_path = snapshot_dir / f"{split_name}.parquet"
if not parquet_path.exists():
raise FileNotFoundError(f"Parquet file not found: {parquet_path}")
row_list = cls._read_parquet_rows(parquet_path)
logger.info("snapshot_store_rows_loaded", split=split_name, count=len(row_list))
return await cls._build_init_kwargs_from_rows(
row_list=row_list,
oracle_url=oracle_url,
db_dsn=db_dsn,
scoring_fn=scoring_fn,
)
@classmethod
async def _async_load_merged(
cls,
snapshot_dir: Path,
split_names: list[str],
oracle_url: str,
db_dsn: str,
scoring_fn: Any | None = None,
) -> dict[str, Any]:
row_list: list[dict[str, Any]] = []
loaded_splits: list[str] = []
for split_name in split_names:
parquet_path = snapshot_dir / f"{split_name}.parquet"
if not parquet_path.exists():
continue
loaded_splits.append(split_name)
row_list.extend(cls._read_parquet_rows(parquet_path))
if not loaded_splits:
raise FileNotFoundError("No snapshot splits found.")
logger.info(
"snapshot_store_rows_loaded",
split="merged",
loaded_splits=loaded_splits,
count=len(row_list),
)
return await cls._build_init_kwargs_from_rows(
row_list=row_list,
oracle_url=oracle_url,
db_dsn=db_dsn,
scoring_fn=scoring_fn,
)
@classmethod
async def _async_materialize_snapshot_dir(
cls,
snapshot_dir: Path,
*,
oracle_url: str,
db_dsn: str,
split_names: list[str] | None = None,
output_dir: Path | None = None,
) -> list[str]:
split_names = list(split_names or ["train", "valid", "test"])
output_dir = output_dir or snapshot_dir
output_dir.mkdir(parents=True, exist_ok=True)
tables_by_split: dict[str, pa.Table] = {}
rows_by_split: dict[str, list[dict[str, Any]]] = {}
row_list: list[dict[str, Any]] = []
needed_columns: set[str] = set()
for split_name in split_names:
parquet_path = snapshot_dir / f"{split_name}.parquet"
if not parquet_path.exists():
continue
table = pq.read_table(str(parquet_path))
rows = table.to_pylist()
tables_by_split[split_name] = table
rows_by_split[split_name] = rows
row_list.extend(rows)
for column in (*_PRICE_DERIVED_FEATURE_COLUMNS, *_MACRO_BACKFILL_FEATURE_COLUMNS):
if column not in table.column_names or table.column(column).null_count > 0:
needed_columns.add(column)
if not tables_by_split:
raise FileNotFoundError(f"No snapshot splits found under {snapshot_dir}")
if not needed_columns:
if output_dir != snapshot_dir:
for split_name in tables_by_split:
pq.write_table(tables_by_split[split_name], output_dir / f"{split_name}.parquet")
manifest_path = snapshot_dir / "manifest.json"
if manifest_path.exists():
(output_dir / "manifest.json").write_text(manifest_path.read_text())
return []
date_range = cls._compute_date_range(row_list)
unique_symbols = sorted({
str(symbol).upper()
for symbol in (
row.get("ticker") or row.get("symbol")
for row in row_list
)
if symbol
})
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]] = {}
macro_by_date: dict[dt.date, dict[str, Any]] = {}
if any(column in needed_columns for column in _PRICE_DERIVED_FEATURE_COLUMNS):
bars_by_symbol, _ = await cls._fetch_price_data(unique_symbols, date_range, oracle_url)
if any(column in needed_columns for column in _MACRO_BACKFILL_FEATURE_COLUMNS):
macro_by_date = await cls._fetch_macro(date_range, db_dsn)
_vix_pairs = sorted(
(d, v["VIXCLS"])
for d, v in macro_by_date.items()
if "VIXCLS" in v and v["VIXCLS"] is not None
)
_hy_series_key = next(
(k for k in next(iter(macro_by_date.values()), {}) if "BAMLH0" in k),
None,
) if macro_by_date else None
_hy_pairs = sorted(
(d, v[_hy_series_key])
for d, v in macro_by_date.items()
if _hy_series_key and _hy_series_key in v and v[_hy_series_key] is not None
) if _hy_series_key else []
_t10y2y_pairs = sorted(
(d, v["T10Y2Y"])
for d, v in macro_by_date.items()
if "T10Y2Y" in v and v["T10Y2Y"] is not None
)
price_bar_cache: dict[str, list[Any]] = {}
derived_feature_cache: dict[tuple[str, dt.date], dict[str, Any]] = {}
materialized_columns: set[str] = set()
for split_name, table in tables_by_split.items():
rows = rows_by_split[split_name]
for row in rows:
symbol = (row.get("ticker") or row.get("symbol") or "")
look_date = (
cls._normalize_date(row.get("event_date"))
or cls._normalize_date(row.get("execution_date"))
or cls._normalize_date(row.get("entry_date"))
)
if symbol and look_date:
if any(column in needed_columns for column in _PRICE_DERIVED_FEATURE_COLUMNS):
cls._backfill_price_derived_features(
row,
symbol=str(symbol).upper(),
event_date=look_date,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
derived_feature_cache=derived_feature_cache,
)
if "macro_vix" in needed_columns and row.get("macro_vix") is None:
value = cls._lookup_as_of(_vix_pairs, look_date)
if value is not None:
row["macro_vix"] = value
if "macro_hy_spread" in needed_columns and row.get("macro_hy_spread") is None:
value = cls._lookup_as_of(_hy_pairs, look_date)
if value is not None:
row["macro_hy_spread"] = value
if "macro_t10y2y" in needed_columns and row.get("macro_t10y2y") is None:
value = cls._lookup_as_of(_t10y2y_pairs, look_date)
if value is not None:
row["macro_t10y2y"] = value
updated_table = table
for column in sorted(needed_columns):
values = [row.get(column) for row in rows]
if not any(value is not None for value in values):
continue
array = pa.array(values, type=pa.float64())
if column in updated_table.column_names:
idx = updated_table.column_names.index(column)
updated_table = updated_table.set_column(idx, column, array)
else:
updated_table = updated_table.append_column(column, array)
materialized_columns.add(column)
pq.write_table(updated_table, output_dir / f"{split_name}.parquet")
manifest_path = snapshot_dir / "manifest.json"
if manifest_path.exists():
manifest = json.loads(manifest_path.read_text())
manifest["output_dir"] = str(output_dir.resolve())
manifest["materialized_feature_columns"] = sorted(materialized_columns)
manifest["materialized_feature_last_refresh_utc"] = dt.datetime.now(dt.UTC).isoformat()
(output_dir / "manifest.json").write_text(json.dumps(manifest, indent=2))
return sorted(materialized_columns)
@staticmethod
def _read_parquet_rows(parquet_path: Path) -> list[dict[str, Any]]:
logger.info("snapshot_store_reading_parquet", path=str(parquet_path))
table = pq.read_table(str(parquet_path))
cols = table.to_pydict()
num_rows = table.num_rows
col_names = list(cols.keys())
return [{col: cols[col][i] for col in col_names} for i in range(num_rows)]
@classmethod
async def _build_init_kwargs_from_rows(
cls,
*,
row_list: list[dict[str, Any]],
oracle_url: str,
db_dsn: str,
scoring_fn: Any | None = None,
) -> dict[str, Any]:
raw_row_count = len(row_list)
# De-duplicate merged-boundary overlaps before any expensive fetches.
deduped_rows: list[dict[str, Any]] = []
seen_keys: set[tuple[Any, ...]] = set()
for row in row_list:
key = (
row.get("event_id"),
row.get("ticker") or row.get("symbol"),
row.get("entry_date") or row.get("execution_date"),
row.get("reaction_date"),
)
if key in seen_keys:
continue
seen_keys.add(key)
deduped_rows.append(row)
row_list = deduped_rows
logger.info(
"snapshot_store_prepare_rows",
raw_rows=raw_row_count,
deduped_rows=len(row_list),
)
# Collect event_ids for DB lookup
event_ids = list(dict.fromkeys(
str(r.get("event_id", ""))
for r in row_list
if r.get("event_id")
))
date_range = cls._compute_date_range(row_list)
# Step 26: DB + Oracle enrichment
logger.info(
"snapshot_store_event_metadata_fetch_start",
event_count=len(event_ids),
)
event_meta = await cls._fetch_event_metadata(event_ids, db_dsn)
logger.info(
"snapshot_store_event_metadata_fetch_done",
resolved_events=len(event_meta),
)
unique_symbols = sorted({
str(symbol).upper()
for symbol in (
[m.get("ticker", "") for m in event_meta.values()]
+ [r.get("ticker", "") for r in row_list]
+ [r.get("symbol", "") for r in row_list]
)
if symbol
})
logger.info(
"snapshot_store_sector_fetch_start",
symbol_count=len(unique_symbols),
)
sectors = await cls._fetch_sectors(unique_symbols, oracle_url)
logger.info(
"snapshot_store_sector_fetch_done",
symbol_count=len(sectors),
)
event_symbols_by_exec_date: dict[dt.date, set[str]] = {}
for row in row_list:
raw_exec = row.get("entry_date") or row.get("execution_date")
exec_date = cls._normalize_date(raw_exec)
if exec_date is None:
continue
event_id = str(row.get("event_id", ""))
meta = event_meta.get(event_id, {})
ticker = (
row.get("ticker")
or row.get("symbol")
or meta.get("ticker")
)
if not ticker:
continue
event_symbols_by_exec_date.setdefault(exec_date, set()).add(str(ticker).upper())
proxy_symbols = sorted(
{
proxy_symbol
for proxy_symbol in (
sector_etf_for_sector(sectors.get(symbol))
for symbol in unique_symbols
)
if proxy_symbol
}
)
peer_proxy_symbols = sorted(
{
proxy_symbol
for symbol in unique_symbols
for proxy_symbol in peer_candidates_for_symbol(symbol, sectors.get(symbol))
if proxy_symbol
}
)
price_symbols = sorted(set(unique_symbols) | set(proxy_symbols) | set(peer_proxy_symbols))
logger.info(
"snapshot_store_enrichment_plan",
symbol_count=len(unique_symbols),
trade_symbol_count=len(price_symbols),
start_date=date_range[0].isoformat() if date_range else None,
end_date=date_range[1].isoformat() if date_range else None,
)
logger.info(
"snapshot_store_price_fetch_start",
symbol_count=len(price_symbols),
)
bars_by_symbol, avg_dvol = await cls._fetch_price_data(
price_symbols, date_range, oracle_url
)
logger.info(
"snapshot_store_price_fetch_done",
symbol_count=len(bars_by_symbol),
avg_dvol_symbols=sum(1 for value in avg_dvol.values() if value > 0),
)
logger.info("snapshot_store_macro_fetch_start")
macro_by_date = await cls._fetch_macro(date_range, db_dsn)
logger.info(
"snapshot_store_macro_fetch_done",
macro_dates=len(macro_by_date),
)
# Fetch SPY bars for macro regime filter (SMA computation)
logger.info("snapshot_store_spy_macro_fetch_start")
spy_macro = await cls._fetch_spy_macro(date_range, oracle_url)
for d, spy_data in spy_macro.items():
macro_by_date.setdefault(d, {}).update(spy_data)
logger.info(
"snapshot_store_spy_macro_fetch_done",
spy_dates=len(spy_macro),
)
# Compute VIX change rates for composite parking gate
sorted_macro_dates = sorted(macro_by_date.keys())
for i, d in enumerate(sorted_macro_dates):
vix = macro_by_date[d].get("VIXCLS")
if vix is not None:
# 5-day VIX change
for lookback in (5, 10):
if i >= lookback:
prev_d = sorted_macro_dates[i - lookback]
prev_vix = macro_by_date.get(prev_d, {}).get("VIXCLS")
if prev_vix is not None:
macro_by_date[d][f"vix_change_{lookback}d"] = vix - prev_vix
# Step 7: Build candidates_by_exec_date
# Pre-build sorted FRED lookup for macro features that may be absent in older snapshots
_vix_pairs = sorted(
(d, v["VIXCLS"])
for d, v in macro_by_date.items()
if "VIXCLS" in v and v["VIXCLS"] is not None
)
_hy_series_key = next(
(k for k in next(iter(macro_by_date.values()), {}) if "BAMLH0" in k),
None,
) if macro_by_date else None
_hy_pairs = sorted(
(d, v[_hy_series_key])
for d, v in macro_by_date.items()
if _hy_series_key and _hy_series_key in v and v[_hy_series_key] is not None
) if _hy_series_key else []
price_bar_cache: dict[str, list[Any]] = {}
derived_feature_cache: dict[tuple[str, dt.date], dict[str, Any]] = {}
def _lookup_fred_as_of(pairs: list, as_of: dt.date) -> float | None:
"""Return most recent FRED value on or before as_of date."""
lo, hi = 0, len(pairs) - 1
result = None
while lo <= hi:
mid = (lo + hi) // 2
if pairs[mid][0] <= as_of:
result = pairs[mid][1]
lo = mid + 1
else:
hi = mid - 1
return result
candidates_by_exec_date: dict[dt.date, list[dict[str, Any]]] = {}
total_rows = len(row_list)
progress_step = max(1000, total_rows // 5) if total_rows else 1000
for idx, row in enumerate(row_list, start=1):
eid = str(row.get("event_id", ""))
meta = event_meta.get(eid, {})
ticker = row.get("ticker") or row.get("symbol") or meta.get("ticker")
if not ticker:
logger.debug("snapshot_store_skip_no_ticker", event_id=eid)
continue
# Map entry_date → execution_date at this boundary
raw_exec = row.get("entry_date") or row.get("execution_date")
if raw_exec is None:
continue
if isinstance(raw_exec, str):
exec_date = dt.date.fromisoformat(raw_exec)
elif isinstance(raw_exec, dt.date):
exec_date = raw_exec
else:
continue
enriched = dict(row)
enriched["execution_date"] = exec_date
enriched["symbol"] = ticker
enriched["issuer_id"] = meta.get("issuer_id") or row.get("issuer_id")
snapshot_event_date = cls._normalize_date(row.get("event_date"))
meta_event_date = cls._normalize_date(meta.get("event_date"))
event_date = snapshot_event_date or meta_event_date
enriched["event_date"] = event_date
snapshot_event_type = row.get("event_type")
enriched["event_type"] = (
snapshot_event_type
if snapshot_event_type not in (None, "")
else meta.get("event_type", "")
)
enriched["event_timestamp"] = (
cls._normalize_timestamp(row.get("event_timestamp"), event_date)
or meta.get("event_timestamp")
or cls._normalize_timestamp(None, event_date)
)
enriched["sector"] = sectors.get(ticker, "UNKNOWN")
# Backfill macro features from FRED data when absent in Parquet
# (older or OOT snapshots may lack these columns)
look_date = enriched["event_date"] or exec_date
if isinstance(look_date, dt.datetime):
look_date = look_date.date()
if look_date is not None:
if enriched.get("macro_vix") is None and _vix_pairs:
enriched["macro_vix"] = _lookup_fred_as_of(_vix_pairs, look_date)
if enriched.get("macro_hy_spread") is None and _hy_pairs:
enriched["macro_hy_spread"] = _lookup_fred_as_of(_hy_pairs, look_date)
if look_date is not None:
cls._backfill_price_derived_features(
enriched,
symbol=ticker,
event_date=look_date,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
derived_feature_cache=derived_feature_cache,
)
cls._backfill_avg_dollar_volume_features(
enriched,
symbol=ticker,
event_date=look_date,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
fallback_avg_dvol=avg_dvol.get(ticker, 0.0),
)
else:
cls._backfill_avg_dollar_volume_features(
enriched,
symbol=ticker,
event_date=None,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
fallback_avg_dvol=avg_dvol.get(ticker, 0.0),
)
reaction_date = cls._normalize_date(enriched.get("reaction_date")) or exec_date
cls._attach_sector_etf_proxy_features(
enriched,
sector=enriched.get("sector"),
reaction_date=reaction_date,
execution_date=exec_date,
bars_by_symbol=bars_by_symbol,
avg_dvol=avg_dvol,
price_bar_cache=price_bar_cache,
)
cls._attach_peer_proxy_features(
enriched,
source_symbol=ticker,
sector=enriched.get("sector"),
reaction_date=reaction_date,
execution_date=exec_date,
event_symbols_by_exec_date=event_symbols_by_exec_date,
bars_by_symbol=bars_by_symbol,
avg_dvol=avg_dvol,
price_bar_cache=price_bar_cache,
)
# Map Parquet-specific columns to canonical backtest names
# event_close (reaction-day close) → entry_price_est baseline
if "entry_price_est" not in enriched and "event_close" in enriched:
enriched["entry_price_est"] = enriched["event_close"]
# score: use existing column or compute from market features
if "score" not in enriched or enriched.get("score") is None:
if scoring_fn is not None:
enriched["score"] = scoring_fn(enriched)
else:
from libs.backtest.scoring import compute_entry_score
enriched["score"] = compute_entry_score(enriched)
candidates_by_exec_date.setdefault(exec_date, []).append(enriched)
if idx % progress_step == 0 or idx == total_rows:
logger.info(
"snapshot_store_candidate_build_progress",
processed=idx,
total=total_rows,
)
cls._attach_recent_sector_cluster_features(candidates_by_exec_date)
# Build bars_by_symbol_date: symbol -> date -> bar dict
bars_by_symbol_date: dict[str, dict[dt.date, dict[str, Any]]] = {}
for sym, date_bars in bars_by_symbol.items():
bars_by_symbol_date[sym] = date_bars
cls._inject_macro_trade_bars(bars_by_symbol_date, macro_by_date)
logger.info(
"snapshot_store_built",
exec_dates=len(candidates_by_exec_date),
symbols=len(bars_by_symbol_date),
candidates=sum(len(rows) for rows in candidates_by_exec_date.values()),
)
return {
"candidates_by_exec_date": candidates_by_exec_date,
"bars_by_symbol_date": bars_by_symbol_date,
"macro_by_date": macro_by_date,
}
@classmethod
def _attach_recent_sector_cluster_features(
cls,
candidates_by_exec_date: dict[dt.date, list[dict[str, Any]]],
) -> None:
def _coerce_float(value: Any) -> float | None:
try:
out = float(value)
except (TypeError, ValueError):
return None
if math.isnan(out):
return None
return out
ordered_rows: list[dict[str, Any]] = []
for rows in candidates_by_exec_date.values():
ordered_rows.extend(rows)
def _event_anchor_date(row: dict[str, Any]) -> dt.date | None:
return (
cls._normalize_date(row.get("event_date"))
or cls._normalize_date(row.get("reaction_date"))
or cls._normalize_date(row.get("execution_date"))
)
def _event_anchor_timestamp(row: dict[str, Any]) -> dt.datetime:
anchor_date = _event_anchor_date(row) or dt.date.max
return (
cls._normalize_timestamp(row.get("event_timestamp"), anchor_date)
or dt.datetime.combine(anchor_date, dt.time.max, tzinfo=dt.UTC)
)
ordered_rows.sort(
key=lambda row: (
_event_anchor_date(row) or dt.date.max,
_event_anchor_timestamp(row),
str(row.get("event_id") or ""),
)
)
history_by_sector: dict[str, list[dict[str, float | dt.date | bool]]] = {}
lookback_days = dt.timedelta(days=3)
for row in ordered_rows:
sector = str(row.get("sector") or "UNKNOWN")
anchor_date = _event_anchor_date(row)
if sector == "UNKNOWN" or anchor_date is None:
row["sector_recent_event_count_3d"] = 0.0
row["sector_recent_leader_count_3d"] = 0.0
row["sector_recent_leader_reaction_max_3d"] = None
row["sector_recent_leader_market_cap_max_3d"] = None
continue
sector_history = history_by_sector.setdefault(sector, [])
cutoff = anchor_date - lookback_days
sector_history[:] = [
item for item in sector_history if isinstance(item["anchor_date"], dt.date) and item["anchor_date"] >= cutoff
]
leaders = [item for item in sector_history if bool(item.get("is_leader"))]
row["sector_recent_event_count_3d"] = float(len(sector_history))
row["sector_recent_leader_count_3d"] = float(len(leaders))
row["sector_recent_leader_reaction_max_3d"] = (
max(float(item["reaction_day_return"]) for item in leaders)
if leaders
else None
)
row["sector_recent_leader_market_cap_max_3d"] = (
max(float(item["market_cap_proxy"]) for item in leaders if item.get("market_cap_proxy") is not None)
if leaders and any(item.get("market_cap_proxy") is not None for item in leaders)
else None
)
reaction_day_return = _coerce_float(row.get("reaction_day_return"))
volume_ratio = _coerce_float(row.get("volume_ratio"))
close_location = _coerce_float(row.get("close_location"))
market_cap_proxy = _coerce_float(row.get("market_cap_proxy"))
is_leader = (
str(row.get("event_type") or "") == "earnings_release"
and reaction_day_return is not None
and reaction_day_return >= 0.08
and volume_ratio is not None
and volume_ratio >= 2.0
and close_location is not None
and close_location >= 0.6
and market_cap_proxy is not None
and market_cap_proxy >= 10_000_000_000.0
)
sector_history.append(
{
"anchor_date": anchor_date,
"reaction_day_return": reaction_day_return or 0.0,
"market_cap_proxy": market_cap_proxy,
"is_leader": is_leader,
}
)
# ------------------------------------------------------------------
# Internal async helpers
# ------------------------------------------------------------------
@staticmethod
def _runtime_cache_dir(snapshot_dir: Path) -> Path:
return snapshot_dir / ".runtime_cache"
@staticmethod
def _scoring_cache_token(scoring_fn: Any | None) -> str:
if scoring_fn is None:
return "none"
if isinstance(scoring_fn, partial):
func = scoring_fn.func
return (
f"partial:{getattr(func, '__module__', '')}.{getattr(func, '__qualname__', getattr(func, '__name__', type(func).__name__))}:"
f"args={repr(scoring_fn.args)}:keywords={repr(scoring_fn.keywords)}"
)
return (
f"callable:{getattr(scoring_fn, '__module__', '')}."
f"{getattr(scoring_fn, '__qualname__', getattr(scoring_fn, '__name__', type(scoring_fn).__name__))}"
)
@classmethod
def _runtime_cache_file(
cls,
snapshot_dir: Path,
split_names: list[str],
scoring_fn: Any | None,
) -> Path:
split_part = "__".join(sorted(split_names))
score_token = cls._scoring_cache_token(scoring_fn)
token_hash = hashlib.sha1(score_token.encode("utf-8")).hexdigest()[:16]
return cls._runtime_cache_dir(snapshot_dir) / f"{split_part}__{token_hash}.pkl"
@classmethod
def _runtime_cache_lock_file(
cls,
snapshot_dir: Path,
split_names: list[str],
scoring_fn: Any | None,
) -> Path:
cache_file = cls._runtime_cache_file(snapshot_dir, split_names, scoring_fn)
return cache_file.with_suffix(f"{cache_file.suffix}.lock")
@staticmethod
def _runtime_cache_fingerprint(snapshot_dir: Path, split_names: list[str]) -> str:
paths: list[Path] = [snapshot_dir / "manifest.json"]
paths.extend(snapshot_dir / f"{split_name}.parquet" for split_name in sorted(split_names))
parts: list[str] = [f"v={_RUNTIME_CACHE_VERSION}"]
for path in paths:
if not path.exists():
parts.append(f"{path.name}:missing")
continue
stat = path.stat()
parts.append(f"{path.name}:{stat.st_size}:{stat.st_mtime_ns}")
return hashlib.sha1("|".join(parts).encode("utf-8")).hexdigest()
@classmethod
def _try_load_runtime_cache(
cls,
snapshot_dir: Path,
split_names: list[str],
*,
scoring_fn: Any | None,
) -> dict[str, Any] | None:
cache_file = cls._runtime_cache_file(snapshot_dir, split_names, scoring_fn)
if not cache_file.exists():
logger.info(
"snapshot_store_runtime_cache_miss",
reason="file_missing",
cache_file=str(cache_file),
)
return None
try:
payload = pickle.loads(cache_file.read_bytes())
expected_fingerprint = cls._runtime_cache_fingerprint(snapshot_dir, split_names)
if payload.get("fingerprint") != expected_fingerprint:
logger.info(
"snapshot_store_runtime_cache_miss",
reason="fingerprint_mismatch",
cache_file=str(cache_file),
)
return None
data = payload.get("data")
if not isinstance(data, dict):
logger.info(
"snapshot_store_runtime_cache_miss",
reason="invalid_payload",
cache_file=str(cache_file),
)
return None
logger.info(
"snapshot_store_runtime_cache_hit",
cache_file=str(cache_file),
split="merged" if len(split_names) > 1 else split_names[0],
)
return data
except Exception as exc:
logger.warning(
"snapshot_store_runtime_cache_read_failed",
cache_file=str(cache_file),
error=str(exc),
)
return None
@classmethod
def _load_with_runtime_cache(
cls,
*,
snapshot_path: Path,
split_names: list[str],
scoring_fn: Any | None,
builder: Any,
) -> "SnapshotStore":
cache_data = cls._try_load_runtime_cache(
snapshot_path,
split_names,
scoring_fn=scoring_fn,
)
if cache_data is not None:
return cls(**cache_data)
lock_file = cls._runtime_cache_lock_file(snapshot_path, split_names, scoring_fn)
if cls._acquire_runtime_cache_lock(lock_file):
try:
cache_data = cls._try_load_runtime_cache(
snapshot_path,
split_names,
scoring_fn=scoring_fn,
)
if cache_data is not None:
return cls(**cache_data)
data = builder()
cls._write_runtime_cache(
snapshot_path,
split_names,
scoring_fn=scoring_fn,
data=data,
)
return cls(**data)
finally:
cls._release_runtime_cache_lock(lock_file)
logger.info(
"snapshot_store_runtime_cache_wait_start",
cache_file=str(cls._runtime_cache_file(snapshot_path, split_names, scoring_fn)),
split="merged" if len(split_names) > 1 else split_names[0],
)
cache_data = cls._wait_for_runtime_cache(
snapshot_path,
split_names,
scoring_fn=scoring_fn,
lock_file=lock_file,
)
if cache_data is not None:
return cls(**cache_data)
if cls._acquire_runtime_cache_lock(lock_file):
try:
cache_data = cls._try_load_runtime_cache(
snapshot_path,
split_names,
scoring_fn=scoring_fn,
)
if cache_data is not None:
return cls(**cache_data)
data = builder()
cls._write_runtime_cache(
snapshot_path,
split_names,
scoring_fn=scoring_fn,
data=data,
)
return cls(**data)
finally:
cls._release_runtime_cache_lock(lock_file)
logger.warning(
"snapshot_store_runtime_cache_wait_failed",
cache_file=str(cls._runtime_cache_file(snapshot_path, split_names, scoring_fn)),
)
data = builder()
cls._write_runtime_cache(
snapshot_path,
split_names,
scoring_fn=scoring_fn,
data=data,
)
return cls(**data)
@staticmethod
def _runtime_cache_lock_payload() -> dict[str, Any]:
return {
"pid": os.getpid(),
"created_at": time.time(),
}
@staticmethod
def _process_is_running(pid: int | None) -> bool:
if not isinstance(pid, int) or pid <= 0:
return False
try:
os.kill(pid, 0)
except ProcessLookupError:
return False
except PermissionError:
return True
return True
@classmethod
def _lock_is_stale(cls, lock_file: Path) -> bool:
try:
payload = json.loads(lock_file.read_text())
except Exception:
return True
pid = payload.get("pid")
created_at = payload.get("created_at")
if not cls._process_is_running(pid):
return True
if not isinstance(created_at, (int, float)):
return True
return (time.time() - float(created_at)) > _RUNTIME_CACHE_WAIT_TIMEOUT_SECONDS
@classmethod
def _acquire_runtime_cache_lock(cls, lock_file: Path) -> bool:
lock_file.parent.mkdir(parents=True, exist_ok=True)
payload = json.dumps(cls._runtime_cache_lock_payload())
for _ in range(2):
try:
fd = os.open(str(lock_file), os.O_CREAT | os.O_EXCL | os.O_WRONLY)
with os.fdopen(fd, "w", encoding="utf-8") as fh:
fh.write(payload)
return True
except FileExistsError:
if cls._lock_is_stale(lock_file):
try:
lock_file.unlink()
except FileNotFoundError:
pass
continue
return False
return False
@staticmethod
def _release_runtime_cache_lock(lock_file: Path) -> None:
try:
lock_file.unlink()
except FileNotFoundError:
pass
@classmethod
def _wait_for_runtime_cache(
cls,
snapshot_dir: Path,
split_names: list[str],
*,
scoring_fn: Any | None,
lock_file: Path,
) -> dict[str, Any] | None:
deadline = time.monotonic() + _RUNTIME_CACHE_WAIT_TIMEOUT_SECONDS
next_log = time.monotonic()
while time.monotonic() < deadline:
cache_file = cls._runtime_cache_file(snapshot_dir, split_names, scoring_fn)
if cache_file.exists():
cache_data = cls._try_load_runtime_cache(
snapshot_dir,
split_names,
scoring_fn=scoring_fn,
)
if cache_data is not None:
return cache_data
if not lock_file.exists():
return None
if cls._lock_is_stale(lock_file):
cls._release_runtime_cache_lock(lock_file)
return None
now = time.monotonic()
if now >= next_log:
logger.info(
"snapshot_store_runtime_cache_waiting",
cache_file=str(cache_file),
split="merged" if len(split_names) > 1 else split_names[0],
)
next_log = now + _RUNTIME_CACHE_WAIT_LOG_INTERVAL_SECONDS
time.sleep(1.0)
return None
@classmethod
def _write_runtime_cache(
cls,
snapshot_dir: Path,
split_names: list[str],
*,
scoring_fn: Any | None,
data: dict[str, Any],
) -> None:
cache_file = cls._runtime_cache_file(snapshot_dir, split_names, scoring_fn)
cache_dir = cache_file.parent
cache_dir.mkdir(parents=True, exist_ok=True)
payload = {
"fingerprint": cls._runtime_cache_fingerprint(snapshot_dir, split_names),
"data": data,
}
tmp_file = cache_file.with_suffix(f"{cache_file.suffix}.tmp.{os.getpid()}")
tmp_file.write_bytes(pickle.dumps(payload, protocol=pickle.HIGHEST_PROTOCOL))
os.replace(tmp_file, cache_file)
logger.info(
"snapshot_store_runtime_cache_written",
cache_file=str(cache_file),
split="merged" if len(split_names) > 1 else split_names[0],
)
@staticmethod
async def _fetch_event_metadata(
event_ids: list[str],
db_dsn: str,
) -> dict[str, dict[str, Any]]:
"""Batch-query Event + SymbolMaster for event metadata."""
if not event_ids:
return {}
try:
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from libs.db.models import Event, SymbolMaster
engine = create_async_engine(db_dsn, echo=False)
async_session = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
result: dict[str, dict[str, Any]] = {}
async with async_session() as session:
_UTC = __import__("zoneinfo").ZoneInfo("UTC")
unique_event_ids = list(dict.fromkeys(eid for eid in event_ids if eid))
for start in range(0, len(unique_event_ids), _EVENT_METADATA_FETCH_CHUNK_SIZE):
chunk = unique_event_ids[start : start + _EVENT_METADATA_FETCH_CHUNK_SIZE]
stmt = (
select(Event, SymbolMaster)
.outerjoin(SymbolMaster, Event.symbol_id == SymbolMaster.symbol_id)
.where(Event.event_id.in_(chunk))
)
rows = (await session.execute(stmt)).all()
for event, sym in rows:
# Use filed_at_utc if available; fallback to filing_date + 21:00 UTC
# (transparent enrichment in the loader — not silent substitution in selector)
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
)
result[event.event_id] = {
"issuer_id": event.issuer_id,
"event_date": event.event_date,
"event_type": event.event_type,
"event_timestamp": ts,
"ticker": sym.ticker if sym else None,
}
await engine.dispose()
return result
except Exception as exc:
logger.warning("snapshot_store_db_fetch_failed", error=str(exc))
return {}
@staticmethod
async def _fetch_price_data(
symbols: list[str],
date_range: tuple[dt.date, dt.date] | None,
oracle_url: str,
concurrency: int = 12,
) -> tuple[dict[str, dict[dt.date, dict[str, Any]]], dict[str, float]]:
"""Fetch daily OHLCV bars and compute avg_dollar_volume per symbol."""
if not symbols or date_range is None:
return {}, {}
try:
from libs.oracle_client import OracleClient, PriceService
fetch_start = date_range[0] - dt.timedelta(days=_PRICE_FEATURE_WARMUP_DAYS)
start_str = fetch_start.isoformat()
end_str = date_range[1].isoformat()
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]] = {}
avg_dvol: dict[str, float] = {}
semaphore = asyncio.Semaphore(concurrency)
total_symbols = len(symbols)
progress_step = max(25, total_symbols // 10) if total_symbols else 25
async with OracleClient(base_url=oracle_url) as client:
svc = PriceService(client)
async def _fetch_symbol(sym: str) -> tuple[str, dict[dt.date, dict[str, Any]], float]:
async with semaphore:
try:
resp = await svc.get_daily_bars(sym, start=start_str, end=end_str)
date_bars: dict[dt.date, dict[str, Any]] = {}
dollar_vols: 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_bars[d] = b
dollar_vols.append(bar.close * bar.volume)
if dollar_vols:
last_20 = dollar_vols[-20:]
mean_dvol = sum(last_20) / len(last_20)
else:
mean_dvol = 0.0
return sym, date_bars, mean_dvol
except Exception as sym_exc:
logger.warning(
"snapshot_store_price_fetch_failed",
symbol=sym,
error=str(sym_exc),
)
return sym, {}, 0.0
tasks = [asyncio.create_task(_fetch_symbol(sym)) for sym in symbols]
completed = 0
for result in asyncio.as_completed(tasks):
sym, date_bars, mean_dvol = await result
bars_by_symbol[sym] = date_bars
avg_dvol[sym] = mean_dvol
completed += 1
if completed % progress_step == 0 or completed == total_symbols:
logger.info(
"snapshot_store_price_fetch_progress",
completed=completed,
total=total_symbols,
)
return bars_by_symbol, avg_dvol
except Exception as exc:
logger.warning("snapshot_store_oracle_failed", error=str(exc))
return {}, {}
@staticmethod
def _backfill_price_derived_features(
row: dict[str, Any],
*,
symbol: str,
event_date: dt.date,
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]],
price_bar_cache: dict[str, list[Any]],
derived_feature_cache: dict[tuple[str, dt.date], dict[str, Any]],
) -> None:
if all(row.get(col) is not None for col in _PRICE_DERIVED_FEATURE_COLUMNS):
return
cache_key = (symbol, event_date)
if cache_key not in derived_feature_cache:
price_bars = SnapshotStore._get_price_bars_for_symbol(
symbol,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
)
if not price_bars:
derived_feature_cache[cache_key] = {}
return
features = compute_market_features(price_bars, event_date.isoformat())
derived_feature_cache[cache_key] = {
col: features.get(col)
for col in _PRICE_DERIVED_FEATURE_COLUMNS
}
for col, value in derived_feature_cache[cache_key].items():
if row.get(col) is None and value is not None:
row[col] = value
@classmethod
def _backfill_avg_dollar_volume_features(
cls,
row: dict[str, Any],
*,
symbol: str,
event_date: dt.date | None,
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]],
price_bar_cache: dict[str, list[Any]],
fallback_avg_dvol: float,
) -> None:
existing_adv_20d = row.get("avg_dollar_volume_20d")
existing_adv = row.get("avg_dollar_volume")
if existing_adv_20d is not None and existing_adv is not None:
return
derived_adv = None
if event_date is not None:
price_bars = cls._get_price_bars_for_symbol(
symbol,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
)
if price_bars:
features = compute_market_features(price_bars, event_date.isoformat())
derived_adv = features.get("avg_dollar_volume_20d")
resolved_adv = (
existing_adv_20d
if existing_adv_20d is not None
else derived_adv
if derived_adv is not None
else existing_adv
if existing_adv is not None
else fallback_avg_dvol
)
if row.get("avg_dollar_volume_20d") is None:
row["avg_dollar_volume_20d"] = resolved_adv
if row.get("avg_dollar_volume") is None:
row["avg_dollar_volume"] = resolved_adv
@staticmethod
def _get_price_bars_for_symbol(
symbol: str,
*,
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]],
price_bar_cache: dict[str, list[Any]],
) -> list[Any]:
if symbol in price_bar_cache:
return price_bar_cache[symbol]
symbol_bars = bars_by_symbol.get(symbol) or {}
if not symbol_bars:
price_bar_cache[symbol] = []
return price_bar_cache[symbol]
from libs.oracle_client.models import PriceBar
price_bar_cache[symbol] = [
PriceBar(
date=bar_date.isoformat(),
open=float(bar.get("open", 0.0) or 0.0),
high=float(bar.get("high", 0.0) or 0.0),
low=float(bar.get("low", 0.0) or 0.0),
close=float(bar.get("close", 0.0) or 0.0),
volume=int(bar.get("volume", 0) or 0),
)
for bar_date, bar in sorted(symbol_bars.items())
]
return price_bar_cache[symbol]
@staticmethod
def _inject_macro_trade_bars(
bars_by_symbol_date: dict[str, dict[dt.date, dict[str, Any]]],
macro_by_date: dict[dt.date, dict[str, Any]],
) -> None:
"""Expose macro ETFs as normal OHLCV bars for synthetic engines."""
for symbol, prefix in _MACRO_TRADE_BAR_PREFIXES:
injected: dict[dt.date, dict[str, Any]] = {}
for bar_date, macro_vals in macro_by_date.items():
close_value = macro_vals.get(f"{prefix}_close")
if close_value is None or float(close_value) <= 0:
continue
injected[bar_date] = {
"open": macro_vals.get(f"{prefix}_open", close_value),
"high": macro_vals.get(f"{prefix}_high", close_value),
"low": macro_vals.get(f"{prefix}_low", close_value),
"close": close_value,
"volume": macro_vals.get(f"{prefix}_volume", 1_000_000),
}
if not injected:
continue
existing = bars_by_symbol_date.setdefault(symbol, {})
for bar_date, bar in injected.items():
existing.setdefault(bar_date, bar)
@classmethod
def _attach_sector_etf_proxy_features(
cls,
row: dict[str, Any],
*,
sector: str | None,
reaction_date: dt.date,
execution_date: dt.date,
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]],
avg_dvol: dict[str, float],
price_bar_cache: dict[str, list[Any]],
) -> None:
proxy_symbol = sector_etf_for_sector(sector)
if not proxy_symbol:
return
row["sector_etf_proxy"] = proxy_symbol
proxy_bars = cls._get_price_bars_for_symbol(
proxy_symbol,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
)
if not proxy_bars:
return
proxy_features = compute_market_features(proxy_bars, reaction_date.isoformat())
row["sector_etf_event_close"] = proxy_features.get("event_close")
row["sector_etf_reaction_day_low"] = proxy_features.get("reaction_day_low")
row["sector_etf_reaction_day_high"] = proxy_features.get("reaction_day_high")
row["sector_etf_reaction_day_return"] = proxy_features.get("reaction_day_return")
row["sector_etf_volume_ratio_20d"] = proxy_features.get("volume_ratio_20d")
row["sector_etf_gap_size"] = proxy_features.get("gap_size")
row["sector_etf_close_location"] = proxy_features.get("close_location")
row["sector_etf_avg_dollar_volume"] = (
proxy_features.get("avg_dollar_volume_20d")
or avg_dvol.get(proxy_symbol, 0.0)
)
row["sector_etf_atr_14"] = proxy_features.get("atr_14")
execution_bar = (bars_by_symbol.get(proxy_symbol) or {}).get(execution_date)
if execution_bar is not None:
row["sector_etf_entry_price"] = (
execution_bar.get("open")
or execution_bar.get("close")
)
else:
row["sector_etf_entry_price"] = row.get("sector_etf_event_close")
@classmethod
def _attach_peer_proxy_features(
cls,
row: dict[str, Any],
*,
source_symbol: str,
sector: str | None,
reaction_date: dt.date,
execution_date: dt.date,
event_symbols_by_exec_date: dict[dt.date, set[str]],
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]],
avg_dvol: dict[str, float],
price_bar_cache: dict[str, list[Any]],
) -> None:
blocked_symbols = {
str(symbol).upper()
for symbol in event_symbols_by_exec_date.get(execution_date, set())
if symbol
}
for proxy_symbol in peer_candidates_for_symbol(source_symbol, sector):
if proxy_symbol in blocked_symbols:
continue
proxy_bars = cls._get_price_bars_for_symbol(
proxy_symbol,
bars_by_symbol=bars_by_symbol,
price_bar_cache=price_bar_cache,
)
if not proxy_bars:
continue
proxy_features = compute_market_features(proxy_bars, reaction_date.isoformat())
execution_bar = (bars_by_symbol.get(proxy_symbol) or {}).get(execution_date)
entry_price = None
if execution_bar is not None:
entry_price = execution_bar.get("open") or execution_bar.get("close")
if not entry_price:
entry_price = proxy_features.get("event_close")
if not entry_price:
continue
row["peer_proxy_symbol"] = proxy_symbol
row["peer_proxy_event_close"] = proxy_features.get("event_close")
row["peer_proxy_reaction_day_low"] = proxy_features.get("reaction_day_low")
row["peer_proxy_reaction_day_high"] = proxy_features.get("reaction_day_high")
row["peer_proxy_reaction_day_return"] = proxy_features.get("reaction_day_return")
row["peer_proxy_volume_ratio_20d"] = proxy_features.get("volume_ratio_20d")
row["peer_proxy_gap_size"] = proxy_features.get("gap_size")
row["peer_proxy_close_location"] = proxy_features.get("close_location")
row["peer_proxy_avg_dollar_volume"] = (
proxy_features.get("avg_dollar_volume_20d")
or avg_dvol.get(proxy_symbol, 0.0)
)
row["peer_proxy_atr_14"] = proxy_features.get("atr_14")
row["peer_proxy_entry_price"] = entry_price
break
@staticmethod
def _lookup_as_of(pairs: list[tuple[dt.date, float]], as_of: dt.date) -> float | None:
lo, hi = 0, len(pairs) - 1
result = None
while lo <= hi:
mid = (lo + hi) // 2
if pairs[mid][0] <= as_of:
result = pairs[mid][1]
lo = mid + 1
else:
hi = mid - 1
return result
@staticmethod
async def _fetch_sectors(
symbols: list[str],
oracle_url: str,
concurrency: int = 12,
) -> dict[str, str]:
"""Fetch company sector for each symbol. Default 'UNKNOWN' if unavailable."""
if not symbols:
return {}
cache = SnapshotStore._load_sector_cache()
result: dict[str, str] = {}
semaphore = asyncio.Semaphore(concurrency)
total_symbols = len(symbols)
progress_step = max(25, total_symbols // 10) if total_symbols else 25
try:
from libs.oracle_client import CompanyService, OracleClient
async with OracleClient(base_url=oracle_url) as client:
company_svc = CompanyService(client)
async def _fetch_sector(sym: str) -> tuple[str, str]:
async with semaphore:
cached = cache.get(sym)
if isinstance(cached, str) and cached:
return sym, cached
try:
info = await company_svc.get_company(sym)
sector = info.sector or "UNKNOWN"
if SnapshotStore._sector_info_is_placeholder(info):
yf_sector = SnapshotStore._fetch_sector_from_yfinance(sym)
if isinstance(yf_sector, str) and yf_sector and yf_sector != "UNKNOWN":
sector = yf_sector
return sym, sector
except Exception:
return sym, cached or "UNKNOWN"
tasks = [asyncio.create_task(_fetch_sector(sym)) for sym in symbols]
completed = 0
for sector_result in asyncio.as_completed(tasks):
sym, sector = await sector_result
result[sym] = sector
completed += 1
if completed % progress_step == 0 or completed == total_symbols:
logger.info(
"snapshot_store_sector_fetch_progress",
completed=completed,
total=total_symbols,
)
except Exception as exc:
logger.warning("snapshot_store_sector_fetch_failed", error=str(exc))
# Default all remaining to UNKNOWN
for sym in symbols:
result.setdefault(sym, cache.get(sym, "UNKNOWN"))
if result[sym] != "UNKNOWN":
cache[sym] = result[sym]
SnapshotStore._write_sector_cache(cache)
return result
@staticmethod
def _sector_cache_path() -> Path:
settings = get_settings()
return Path(settings.data_root) / "cache" / "sector_cache.json"
@staticmethod
def _load_sector_cache() -> dict[str, str]:
path = SnapshotStore._sector_cache_path()
if not path.exists():
return {}
try:
return json.loads(path.read_text())
except Exception:
return {}
@staticmethod
def _write_sector_cache(cache: dict[str, str]) -> None:
path = SnapshotStore._sector_cache_path()
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(dict(sorted(cache.items())), indent=2))
@staticmethod
def _sector_info_is_placeholder(info: Any) -> bool:
sector = getattr(info, "sector", None)
industry = getattr(info, "industry", None)
exchange = getattr(info, "exchange", None)
market_cap = getattr(info, "market_cap", None)
return (
sector == "Technology"
and industry == "Software"
and exchange is None
and market_cap is None
)
@staticmethod
def _fetch_sector_from_yfinance(symbol: str) -> str:
try:
import yfinance as yf
info = yf.Ticker(symbol).get_info()
sector = info.get("sector")
if isinstance(sector, str) and sector.strip():
return sector.strip()
except Exception:
return "UNKNOWN"
return "UNKNOWN"
@staticmethod
async def _fetch_macro(
date_range: tuple[dt.date, dt.date] | None,
db_dsn: str,
) -> dict[dt.date, dict[str, Any]]:
"""Load MacroObservation regime data from DB."""
if date_range is None:
return {}
try:
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from libs.db.models import MacroObservation
engine = create_async_engine(db_dsn, echo=False)
async_session = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
result: dict[dt.date, dict[str, Any]] = {}
async with async_session() as session:
stmt = select(MacroObservation).where(
MacroObservation.observation_date >= date_range[0],
MacroObservation.observation_date <= date_range[1],
)
rows = (await session.execute(stmt)).scalars().all()
for obs in rows:
d = obs.observation_date
result.setdefault(d, {})[obs.series_id] = (
float(obs.value) if obs.value is not None else None
)
await engine.dispose()
return result
except Exception as exc:
logger.warning("snapshot_store_macro_fetch_failed", error=str(exc))
return {}
@staticmethod
async def _fetch_spy_macro(
date_range: tuple[dt.date, dt.date] | None,
oracle_url: str,
sma_periods: tuple[int, ...] = (10, 20, 30, 40, 50),
) -> dict[dt.date, dict[str, Any]]:
"""Fetch SPY/QQQ daily bars and compute SMAs for macro regime filtering.
Returns dict with benchmark closes/SMAs merged by date.
SMA is None for the first (period - 1) bars of each period.
"""
if date_range is None:
return {}
try:
from libs.oracle_client import OracleClient, PriceService
max_period = max(sma_periods)
# Extend start date back by max SMA period for warm-up
warmup_days = max(max_period, 80) * 2 # calendar days (Hurst needs 60 trading days)
extended_start = date_range[0] - dt.timedelta(days=warmup_days)
async with OracleClient(base_url=oracle_url) as client:
svc = PriceService(client)
spy_resp = await svc.get_daily_bars(
"SPY", start=extended_start.isoformat(), end=date_range[1].isoformat()
)
qqq_resp = await svc.get_daily_bars(
"QQQ", start=extended_start.isoformat(), end=date_range[1].isoformat()
)
extra_symbol_responses: dict[str, Any | None] = {}
for extra_symbol in (
"TQQQ",
"QQQM",
"SPYM",
"QUAL",
"SH",
"PSQ",
"GLD",
"SHY",
"USFR",
"BIL",
"VGSH",
"IEI",
"IEF",
"TIP",
"XLK",
"SMH",
"IWM",
"XLI",
"XLE",
"XLF",
"DBC",
"SGOV",
"JEPQ",
"BUFB",
"MERIX",
):
try:
extra_symbol_responses[extra_symbol] = await svc.get_daily_bars(
extra_symbol,
start=extended_start.isoformat(),
end=date_range[1].isoformat(),
)
except Exception:
extra_symbol_responses[extra_symbol] = None
tqqq_resp = extra_symbol_responses.get("TQQQ")
qqqm_resp = extra_symbol_responses.get("QQQM")
spym_resp = extra_symbol_responses.get("SPYM")
qual_resp = extra_symbol_responses.get("QUAL")
sh_resp = extra_symbol_responses.get("SH")
psq_resp = extra_symbol_responses.get("PSQ")
gld_resp = extra_symbol_responses.get("GLD")
shy_resp = extra_symbol_responses.get("SHY")
usfr_resp = extra_symbol_responses.get("USFR")
bil_resp = extra_symbol_responses.get("BIL")
vgsh_resp = extra_symbol_responses.get("VGSH")
iei_resp = extra_symbol_responses.get("IEI")
ief_resp = extra_symbol_responses.get("IEF")
tip_resp = extra_symbol_responses.get("TIP")
xlk_resp = extra_symbol_responses.get("XLK")
smh_resp = extra_symbol_responses.get("SMH")
iwm_resp = extra_symbol_responses.get("IWM")
xli_resp = extra_symbol_responses.get("XLI")
xle_resp = extra_symbol_responses.get("XLE")
xlf_resp = extra_symbol_responses.get("XLF")
dbc_resp = extra_symbol_responses.get("DBC")
sgov_resp = extra_symbol_responses.get("SGOV")
jepq_resp = extra_symbol_responses.get("JEPQ")
bufb_resp = extra_symbol_responses.get("BUFB")
merix_resp = extra_symbol_responses.get("MERIX")
result: dict[dt.date, dict[str, Any]] = {}
# Lookback periods for rolling high / momentum / volatility
rolling_high_periods = (5, 20, 50, 100)
momentum_periods = (5, 10, 20, 50)
vol_periods = (5, 15, 20, 30, 50)
efficiency_periods = (10, 20)
downside_vol_periods = (10, 20)
ulcer_periods = (10, 20)
drawdown_accel_days = 5
def _merge_series(prefix: str, bars: Any) -> None:
import math
sorted_bars = sorted(bars, key=lambda b: b.date)
closes: list[tuple[dt.date, float]] = [
(dt.date.fromisoformat(b.date), float(b.close)) for b in sorted_bars
]
for i, (d, close) in enumerate(closes):
if d < date_range[0]:
continue
result.setdefault(d, {})
result[d][f"{prefix}_close"] = close
result[d][f"{prefix}_open"] = float(sorted_bars[i].open)
result[d][f"{prefix}_high"] = float(sorted_bars[i].high)
result[d][f"{prefix}_low"] = float(sorted_bars[i].low)
result[d][f"{prefix}_volume"] = float(sorted_bars[i].volume)
for period in sma_periods:
sma = None
if i >= period - 1:
window = [c for _, c in closes[i - period + 1 : i + 1]]
sma = sum(window) / len(window)
result[d][f"{prefix}_sma_{period}"] = sma
# Rolling high (for drawdown gate)
for rh_p in rolling_high_periods:
rh = None
if i >= rh_p - 1:
rh = max(c for _, c in closes[i - rh_p + 1 : i + 1])
result[d][f"{prefix}_high_{rh_p}"] = rh
# Momentum / N-day return (for momentum gate)
for mom_p in momentum_periods:
mom = None
if i >= mom_p:
prev_close = closes[i - mom_p][1]
if prev_close > 0:
mom = (close - prev_close) / prev_close
result[d][f"{prefix}_mom_{mom_p}"] = mom
# Realized volatility (annualized std of daily log returns)
for vol_p in vol_periods:
vol = None
if i >= vol_p:
log_rets = []
for j in range(i - vol_p + 1, i + 1):
if closes[j - 1][1] > 0:
log_rets.append(math.log(closes[j][1] / closes[j - 1][1]))
if len(log_rets) >= vol_p - 1:
mean_r = sum(log_rets) / len(log_rets)
var_r = sum((r - mean_r) ** 2 for r in log_rets) / len(log_rets)
vol = math.sqrt(var_r * 252) # annualized
result[d][f"{prefix}_vol_{vol_p}"] = vol
# Trend efficiency ratio (Kaufman): net progress / total path length.
# High = smooth persistent trend, low = noisy/random walk.
for eff_p in efficiency_periods:
efficiency = None
if i >= eff_p:
net_move = abs(close - closes[i - eff_p][1])
gross_move = sum(
abs(closes[j][1] - closes[j - 1][1])
for j in range(i - eff_p + 1, i + 1)
)
if gross_move > 0:
efficiency = net_move / gross_move
else:
efficiency = 0.0
result[d][f"{prefix}_efficiency_{eff_p}"] = efficiency
# Downside semivolatility: only penalize harmful volatility.
for dv_p in downside_vol_periods:
downside_vol = None
if i >= dv_p:
neg_sq = []
for j in range(i - dv_p + 1, i + 1):
if closes[j - 1][1] > 0:
ret = closes[j][1] / closes[j - 1][1] - 1
neg_sq.append(min(ret, 0.0) ** 2)
if len(neg_sq) >= dv_p - 1:
downside_vol = math.sqrt(sum(neg_sq) / len(neg_sq) * 252)
result[d][f"{prefix}_downside_vol_{dv_p}"] = downside_vol
# Shannon entropy of daily returns (measures market predictability)
# Low entropy = trending (predictable), high entropy = chaotic (uncertain)
for ent_p in (10, 20):
entropy = None
if i >= ent_p:
# Bin daily returns into categories
daily_rets = []
for j in range(i - ent_p + 1, i + 1):
if closes[j - 1][1] > 0:
daily_rets.append(closes[j][1] / closes[j - 1][1] - 1)
if len(daily_rets) >= ent_p - 1:
# Count positive/negative/flat days
n_pos = sum(1 for r in daily_rets if r > 0.001)
n_neg = sum(1 for r in daily_rets if r < -0.001)
n_flat = len(daily_rets) - n_pos - n_neg
n_total = len(daily_rets)
# Shannon entropy H = -sum(p * log2(p))
entropy = 0.0
for count in (n_pos, n_neg, n_flat):
if count > 0:
p = count / n_total
entropy -= p * math.log2(p)
result[d][f"{prefix}_entropy_{ent_p}"] = entropy
# Ulcer index: rolling RMS drawdown pain from recent peaks.
for ulcer_p in ulcer_periods:
ulcer = None
current_dd = None
if i >= ulcer_p - 1:
window_closes = [c for _, c in closes[i - ulcer_p + 1 : i + 1]]
peak_in_window = 0.0
drawdowns = []
for window_close in window_closes:
peak_in_window = max(peak_in_window, window_close)
if peak_in_window > 0:
drawdowns.append(window_close / peak_in_window - 1.0)
if drawdowns:
ulcer = math.sqrt(sum(dd * dd for dd in drawdowns) / len(drawdowns))
current_dd = abs(drawdowns[-1])
result[d][f"{prefix}_ulcer_{ulcer_p}"] = ulcer
result[d][f"{prefix}_drawdown_{ulcer_p}"] = current_dd
# Drawdown acceleration: how quickly recent pain is worsening.
dd_accel = None
dd_lb = 20
if i >= dd_lb - 1 + drawdown_accel_days:
cur_window = [c for _, c in closes[i - dd_lb + 1 : i + 1]]
prev_i = i - drawdown_accel_days
prev_window = [c for _, c in closes[prev_i - dd_lb + 1 : prev_i + 1]]
cur_peak = max(cur_window) if cur_window else 0.0
prev_peak = max(prev_window) if prev_window else 0.0
if cur_peak > 0 and prev_peak > 0:
cur_dd = (cur_peak - close) / cur_peak
prev_close = closes[prev_i][1]
prev_dd = (prev_peak - prev_close) / prev_peak
dd_accel = cur_dd - prev_dd
result[d][f"{prefix}_drawdown_accel_{drawdown_accel_days}"] = dd_accel
# Hurst exponent via R/S analysis (fractal dimension)
# H > 0.5 = trending/persistent, H < 0.5 = mean-reverting, H ≈ 0.5 = random walk
hurst_lookback = 60
hurst = None
if i >= hurst_lookback + 1:
h_rets = []
for j in range(i - hurst_lookback, i):
if closes[j][1] > 0:
h_rets.append((closes[j + 1][1] - closes[j][1]) / closes[j][1])
if len(h_rets) >= 30:
def _rs_stat(series: list[float]) -> float:
n_ = len(series)
mean_ = sum(series) / n_
devs = [x - mean_ for x in series]
cumdev = []
s_ = 0.0
for dd in devs:
s_ += dd
cumdev.append(s_)
r_ = max(cumdev) - min(cumdev)
std_ = (sum(dd ** 2 for dd in devs) / n_) ** 0.5
return r_ / std_ if std_ > 0 else 0.0
win_sizes = [s for s in [8, 12, 16, 24, 32] if s <= len(h_rets) // 2]
if len(win_sizes) >= 2:
log_n, log_rs = [], []
for w in win_sizes:
rs_vals = []
for st in range(0, len(h_rets) - w + 1, w):
chunk = h_rets[st:st + w]
if len(chunk) == w:
rs_vals.append(_rs_stat(chunk))
if rs_vals:
avg_rs = sum(rs_vals) / len(rs_vals)
if avg_rs > 0:
log_n.append(math.log(w))
log_rs.append(math.log(avg_rs))
if len(log_n) >= 2:
n_h = len(log_n)
x_m = sum(log_n) / n_h
y_m = sum(log_rs) / n_h
num = sum((log_n[k] - x_m) * (log_rs[k] - y_m) for k in range(n_h))
den = sum((log_n[k] - x_m) ** 2 for k in range(n_h))
hurst = num / den if den > 0 else 0.5
result[d][f"{prefix}_hurst_60"] = hurst
# Rolling excess kurtosis (fat tail detection)
# Normal = 0, high = extreme moves more likely (Mandelbrot/Taleb)
kurt_lookback = 20
kurtosis = None
if i >= kurt_lookback:
k_rets = []
for j in range(i - kurt_lookback + 1, i + 1):
if closes[j - 1][1] > 0:
k_rets.append(math.log(closes[j][1] / closes[j - 1][1]))
if len(k_rets) >= kurt_lookback - 1:
k_mean = sum(k_rets) / len(k_rets)
k_var = sum((r - k_mean) ** 2 for r in k_rets) / len(k_rets)
if k_var > 1e-12:
m4 = sum((r - k_mean) ** 4 for r in k_rets) / len(k_rets)
kurtosis = m4 / (k_var ** 2) - 3.0
result[d][f"{prefix}_kurtosis_20"] = kurtosis
# Return autocorrelation (lag-1 Pearson correlation)
# Positive = trending, negative = mean-reverting (Lo, 2004)
ac_lookback = 20
autocorr = None
if i >= ac_lookback + 1:
ac_rets = []
for j in range(i - ac_lookback, i + 1):
if closes[j - 1][1] > 0:
ac_rets.append(closes[j][1] / closes[j - 1][1] - 1)
if len(ac_rets) >= ac_lookback:
x_ac = ac_rets[:-1]
y_ac = ac_rets[1:]
n_ac = len(x_ac)
mx = sum(x_ac) / n_ac
my = sum(y_ac) / n_ac
cov_xy = sum((x_ac[k] - mx) * (y_ac[k] - my) for k in range(n_ac)) / n_ac
sx = (sum((x_ac[k] - mx) ** 2 for k in range(n_ac)) / n_ac) ** 0.5
sy = (sum((y_ac[k] - my) ** 2 for k in range(n_ac)) / n_ac) ** 0.5
if sx > 1e-12 and sy > 1e-12:
autocorr = cov_xy / (sx * sy)
result[d][f"{prefix}_autocorr_20"] = autocorr
_merge_series("spy", spy_resp.bars)
_merge_series("qqq", qqq_resp.bars)
if spym_resp and hasattr(spym_resp, "bars") and spym_resp.bars:
_merge_series("spym", spym_resp.bars)
if qual_resp and hasattr(qual_resp, "bars") and qual_resp.bars:
_merge_series("qual", qual_resp.bars)
if xlk_resp and hasattr(xlk_resp, "bars") and xlk_resp.bars:
_merge_series("xlk", xlk_resp.bars)
if smh_resp and hasattr(smh_resp, "bars") and smh_resp.bars:
_merge_series("smh", smh_resp.bars)
if iwm_resp and hasattr(iwm_resp, "bars") and iwm_resp.bars:
_merge_series("iwm", iwm_resp.bars)
if xli_resp and hasattr(xli_resp, "bars") and xli_resp.bars:
_merge_series("xli", xli_resp.bars)
if xle_resp and hasattr(xle_resp, "bars") and xle_resp.bars:
_merge_series("xle", xle_resp.bars)
if xlf_resp and hasattr(xlf_resp, "bars") and xlf_resp.bars:
_merge_series("xlf", xlf_resp.bars)
def _merge_pair_correlation(left_prefix: str, right_prefix: str, output_key: str) -> None:
# Rolling cross-asset correlation (regime shift detection).
corr_lookback = 20
sorted_result_dates = sorted(result.keys())
for idx_c, d_c in enumerate(sorted_result_dates):
corr_val = None
if idx_c >= corr_lookback:
left_r, right_r = [], []
for jj in range(idx_c - corr_lookback + 1, idx_c + 1):
d_j = sorted_result_dates[jj]
d_prev = sorted_result_dates[jj - 1]
lc = result.get(d_j, {}).get(f"{left_prefix}_close")
lp = result.get(d_prev, {}).get(f"{left_prefix}_close")
rc = result.get(d_j, {}).get(f"{right_prefix}_close")
rp = result.get(d_prev, {}).get(f"{right_prefix}_close")
if all(v and v > 0 for v in [lc, lp, rc, rp]):
left_r.append(lc / lp - 1)
right_r.append(rc / rp - 1)
if len(left_r) >= corr_lookback - 2:
import math as _m
n_cr = len(left_r)
mx_l = sum(left_r) / n_cr
mx_r = sum(right_r) / n_cr
cov_lr = sum(
(left_r[k] - mx_l) * (right_r[k] - mx_r) for k in range(n_cr)
) / n_cr
ss_l = _m.sqrt(sum((left_r[k] - mx_l) ** 2 for k in range(n_cr)) / n_cr)
ss_r = _m.sqrt(sum((right_r[k] - mx_r) ** 2 for k in range(n_cr)) / n_cr)
if ss_l > 1e-12 and ss_r > 1e-12:
corr_val = cov_lr / (ss_l * ss_r)
result[d_c][output_key] = corr_val
# SPY/SPYM-vs-QQQ rolling correlation (regime shift detection)
_merge_pair_correlation("spy", "qqq", "spy_qqq_corr_20")
if spym_resp and hasattr(spym_resp, "bars") and spym_resp.bars:
_merge_pair_correlation("spym", "qqq", "spym_qqq_corr_20")
if qual_resp and hasattr(qual_resp, "bars") and qual_resp.bars:
_merge_pair_correlation("qual", "qqq", "qual_qqq_corr_20")
if tqqq_resp and hasattr(tqqq_resp, 'bars') and tqqq_resp.bars:
# TQQQ: close for valuation/exits, open for entry price (market-open buy)
for b in sorted(tqqq_resp.bars, key=lambda b: b.date):
d = dt.date.fromisoformat(b.date)
if d >= date_range[0]:
result.setdefault(d, {})
result[d]["tqqq_close"] = float(b.close)
result[d]["tqqq_open"] = float(b.open)
if qqqm_resp and hasattr(qqqm_resp, 'bars') and qqqm_resp.bars:
# QQQM: close for valuation/exits, open for entry price (market-open buy)
for b in sorted(qqqm_resp.bars, key=lambda b: b.date):
d = dt.date.fromisoformat(b.date)
if d >= date_range[0]:
result.setdefault(d, {})
result[d]["qqqm_close"] = float(b.close)
result[d]["qqqm_open"] = float(b.open)
if sgov_resp and hasattr(sgov_resp, 'bars') and sgov_resp.bars:
# SGOV: close for valuation/exits, open for entry price (market-open buy)
for b in sorted(sgov_resp.bars, key=lambda b: b.date):
d = dt.date.fromisoformat(b.date)
if d >= date_range[0]:
result.setdefault(d, {})
result[d]["sgov_close"] = float(b.close)
result[d]["sgov_open"] = float(b.open)
if sh_resp and hasattr(sh_resp, 'bars') and sh_resp.bars:
# SH: ProShares Short S&P 500 — full indicators for macro_short engine ATR calculation
_merge_series("sh", sh_resp.bars)
if psq_resp and hasattr(psq_resp, 'bars') and psq_resp.bars:
_merge_series("psq", psq_resp.bars)
if gld_resp and hasattr(gld_resp, 'bars') and gld_resp.bars:
_merge_series("gld", gld_resp.bars)
if shy_resp and hasattr(shy_resp, 'bars') and shy_resp.bars:
_merge_series("shy", shy_resp.bars)
if usfr_resp and hasattr(usfr_resp, 'bars') and usfr_resp.bars:
_merge_series("usfr", usfr_resp.bars)
if bil_resp and hasattr(bil_resp, 'bars') and bil_resp.bars:
_merge_series("bil", bil_resp.bars)
if vgsh_resp and hasattr(vgsh_resp, 'bars') and vgsh_resp.bars:
_merge_series("vgsh", vgsh_resp.bars)
if iei_resp and hasattr(iei_resp, 'bars') and iei_resp.bars:
_merge_series("iei", iei_resp.bars)
if ief_resp and hasattr(ief_resp, 'bars') and ief_resp.bars:
_merge_series("ief", ief_resp.bars)
if tip_resp and hasattr(tip_resp, 'bars') and tip_resp.bars:
_merge_series("tip", tip_resp.bars)
if dbc_resp and hasattr(dbc_resp, 'bars') and dbc_resp.bars:
_merge_series("dbc", dbc_resp.bars)
if jepq_resp and hasattr(jepq_resp, 'bars') and jepq_resp.bars:
_merge_series("jepq", jepq_resp.bars)
if bufb_resp and hasattr(bufb_resp, 'bars') and bufb_resp.bars:
_merge_series("bufb", bufb_resp.bars)
if merix_resp and hasattr(merix_resp, 'bars') and merix_resp.bars:
_merge_series("merix", merix_resp.bars)
# Fetch VIX index for macro regime filtering (VIXCLS equivalent)
# Used as fallback when the snapshot Parquet doesn't embed macro_vix
try:
async with OracleClient(base_url=oracle_url) as vix_client:
vix_svc = PriceService(vix_client)
vix_resp = await vix_svc.get_daily_bars(
"^VIX", start=extended_start.isoformat(), end=date_range[1].isoformat()
)
for b in sorted(vix_resp.bars, key=lambda b: b.date):
d = dt.date.fromisoformat(b.date)
if d >= date_range[0]:
result.setdefault(d, {})
result[d]["VIXCLS"] = float(b.close)
except Exception:
pass # VIX data optional — falls back to Parquet-embedded macro_vix
logger.info(
"snapshot_store_spy_macro_loaded",
bars=len(spy_resp.bars),
dates_with_sma=sum(1 for v in result.values() if v.get("spy_sma_20") is not None),
)
return result
except Exception as exc:
logger.warning("snapshot_store_spy_macro_failed", error=str(exc))
return {}
@staticmethod
def _compute_date_range(
rows: list[dict[str, Any]],
) -> tuple[dt.date, dt.date] | None:
"""Compute (min_date, max_date) from execution and reaction-date columns."""
dates: list[dt.date] = []
for r in rows:
for raw in (
r.get("event_date"),
r.get("entry_date"),
r.get("execution_date"),
r.get("reaction_date"),
):
if raw is None:
continue
if isinstance(raw, str):
try:
dates.append(dt.date.fromisoformat(raw))
except ValueError:
pass
elif isinstance(raw, dt.date):
dates.append(raw)
if not dates:
return None
return min(dates), max(dates)
@staticmethod
def _build_reaction_index(
candidates_by_exec_date: dict[dt.date, list[dict[str, Any]]],
) -> dict[dt.date, list[dict[str, Any]]]:
reaction_index: dict[dt.date, list[dict[str, Any]]] = {}
for rows in candidates_by_exec_date.values():
for row in rows:
reaction_date = SnapshotStore._normalize_date(row.get("reaction_date"))
if reaction_date is None:
continue
reaction_index.setdefault(reaction_date, []).append(row)
return reaction_index
@staticmethod
def _normalize_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)
except ValueError:
return None
return None
@staticmethod
def _normalize_timestamp(raw: Any, fallback_event_date: dt.date | None = None) -> dt.datetime | None:
if isinstance(raw, dt.datetime):
return raw if raw.tzinfo is not None else raw.replace(tzinfo=_UTC)
if isinstance(raw, str):
try:
parsed = dt.datetime.fromisoformat(raw)
except ValueError:
parsed = None
if parsed is not None:
return parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=_UTC)
if fallback_event_date is not None:
return dt.datetime.combine(fallback_event_date, dt.time(21, 0), tzinfo=_UTC)
return None