"""Unit tests for event_gen.py.""" import datetime as dt import numpy as np import pytest from libs.backtest.scenarios.event_gen import ( EventDistribution, _compute_synthetic_score, _sample_categorical, generate_events, ) from libs.backtest.scenarios.price_gen import PriceRegime, generate_price_paths _DATES_RAW = [dt.date(2024, 1, 2) + dt.timedelta(days=i) for i in range(400)] _TRADING_DATES = [d for d in _DATES_RAW if d.weekday() < 5][:252] _SYMBOLS = [f"SYM{i:03d}" for i in range(30)] @pytest.mark.unit class TestSampleCategorical: def test_returns_valid_key(self): rng = np.random.default_rng(1) cats = {"a": 0.5, "b": 0.3, "c": 0.2} result = _sample_categorical(cats, rng) assert result in cats def test_distribution_roughly_correct(self): rng = np.random.default_rng(42) cats = {"x": 0.9, "y": 0.1} counts = {"x": 0, "y": 0} for _ in range(1000): k = _sample_categorical(cats, rng) counts[k] += 1 # x should appear roughly 90% of the time assert counts["x"] > 800 @pytest.mark.unit class TestComputeSyntheticScore: def test_high_quality_event_has_high_score(self): row = { "signal_strength_score": 0.95, "document_quality_score": 0.95, "parse_confidence_overall": 0.95, "guidance_direction_score": 0.90, "oneoff_penalty": 0.0, "reaction_day_return": 0.08, "volume_ratio_20d": 3.0, "close_location": 0.85, "gap_size": 0.02, "pre_event_entropy_60d": 1.0, } score = _compute_synthetic_score(row) assert score > 0.75 def test_low_quality_event_has_low_score(self): row = { "signal_strength_score": 0.20, "document_quality_score": 0.20, "parse_confidence_overall": 0.25, "guidance_direction_score": 0.10, "oneoff_penalty": 1.0, "reaction_day_return": -0.05, "volume_ratio_20d": 0.5, "close_location": 0.2, "gap_size": -0.01, "pre_event_entropy_60d": 1.8, } score = _compute_synthetic_score(row) assert score < 0.35 def test_score_bounded_0_to_1(self): """Score must always be in [0, 1].""" rng = np.random.default_rng(7) for _ in range(50): row = { "signal_strength_score": float(rng.uniform(0, 1)), "document_quality_score": float(rng.uniform(0, 1)), "parse_confidence_overall": float(rng.uniform(0, 1)), "guidance_direction_score": float(rng.uniform(0, 1)), "oneoff_penalty": float(rng.choice([0.0, 1.0])), "reaction_day_return": float(rng.normal(0, 0.05)), "volume_ratio_20d": float(rng.uniform(0.3, 4.0)), "pre_event_entropy_60d": float(rng.uniform(0.5, 2.0)), } score = _compute_synthetic_score(row) assert 0.0 <= score <= 1.0, f"score {score} out of [0, 1]" def test_oneoff_penalty_reduces_score(self): base_row = { "signal_strength_score": 0.7, "document_quality_score": 0.7, "parse_confidence_overall": 0.8, "guidance_direction_score": 0.6, "oneoff_penalty": 0.0, "reaction_day_return": 0.03, "volume_ratio_20d": 2.0, "pre_event_entropy_60d": 1.3, } penalized_row = {**base_row, "oneoff_penalty": 1.0} assert _compute_synthetic_score(base_row) > _compute_synthetic_score(penalized_row) @pytest.mark.unit class TestGenerateEvents: def test_returns_dict_keyed_by_exec_date(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) assert isinstance(candidates, dict) for d in candidates: assert isinstance(d, dt.date) assert d in _TRADING_DATES def test_all_exec_dates_within_eligible_range(self): rng = np.random.default_rng(42) dist = EventDistribution() buffer = 25 candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng, max_holding_days_buffer=buffer) # First 5 days are warm-up; last buffer days are excluded eligible_end = _TRADING_DATES[len(_TRADING_DATES) - buffer - 1] for d in candidates: assert d <= eligible_end, f"exec_date {d} beyond eligible range" def test_required_fields_present(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) required_fields = [ "event_id", "symbol", "event_type", "event_direction", "filing_time_bucket", "sector", "event_timestamp", "reaction_date", "entry_date", "execution_date", "entry_price", "entry_price_est", "event_close", "atr_14", "avg_dollar_volume", "score", "signal_strength_score", "document_quality_score", "parse_confidence_overall", "guidance_direction_score", "oneoff_penalty", "reaction_day_return", "volume_ratio_20d", "close_location", "gap_size", "pre_event_rsi_14", "pre_event_bb_position", "pre_event_volatility_20d", "pre_event_hurst_60d", "pre_event_entropy_60d", "pre_event_ou_theta_60d", "pre_event_market_temperature", ] for exec_date, rows in candidates.items(): for row in rows: for field in required_fields: assert field in row, f"Missing field '{field}' in candidate on {exec_date}" def test_execution_date_matches_dict_key(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) for exec_date, rows in candidates.items(): for row in rows: assert row["execution_date"] == exec_date def test_reaction_date_is_previous_trading_day(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) date_set = set(_TRADING_DATES) for exec_date, rows in candidates.items(): exec_idx = _TRADING_DATES.index(exec_date) expected_reaction = _TRADING_DATES[exec_idx - 1] for row in rows: rxn = row["reaction_date"] # reaction_date stored as isoformat string rxn_date = dt.date.fromisoformat(rxn) if isinstance(rxn, str) else rxn assert rxn_date == expected_reaction, f"reaction_date mismatch on {exec_date}" def test_event_id_unique(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) all_ids = [row["event_id"] for rows in candidates.values() for row in rows] assert len(all_ids) == len(set(all_ids)), "Duplicate event_ids found" def test_no_duplicate_symbol_per_day(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) for exec_date, rows in candidates.items(): syms = [r["symbol"] for r in rows] assert len(syms) == len(set(syms)), f"Duplicate symbols on {exec_date}" def test_prices_positive(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) for rows in candidates.values(): for row in rows: assert row["entry_price"] > 0, f"Non-positive entry_price: {row['entry_price']}" assert row["atr_14"] > 0, f"Non-positive atr_14: {row['atr_14']}" def test_score_between_0_and_1(self): rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) for rows in candidates.values(): for row in rows: assert 0.0 <= row["score"] <= 1.0, f"score {row['score']} out of [0, 1]" def test_uses_bar_close_when_bars_provided(self): rng = np.random.default_rng(42) dist = EventDistribution(events_per_day_mean=2.0, events_per_day_std=0.0) regimes = [PriceRegime(0.10, 0.15, 252)] bars_by_symbol = generate_price_paths( n_symbols=len(_SYMBOLS), initial_prices=None, regimes=regimes, trading_dates=_TRADING_DATES, rng=np.random.default_rng(1), tickers=_SYMBOLS, ) candidates = generate_events( _TRADING_DATES, _SYMBOLS, dist, np.random.default_rng(42), bars_by_symbol=bars_by_symbol, ) # At least some events should be present assert len(candidates) > 0 # entry_prices should be positive and reasonable for rows in candidates.values(): for row in rows: assert row["entry_price"] > 0 def test_event_timestamp_format(self): """event_timestamp should be ISO with timezone.""" rng = np.random.default_rng(42) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) for rows in candidates.values(): for row in rows: ts = row["event_timestamp"] assert "T" in ts, f"event_timestamp not ISO: {ts}" assert "+00:00" in ts or "Z" in ts, f"event_timestamp missing TZ: {ts}" def test_produces_events_with_default_distribution(self): """Default distribution should produce at least some events.""" rng = np.random.default_rng(99) dist = EventDistribution() candidates = generate_events(_TRADING_DATES, _SYMBOLS, dist, rng) total = sum(len(v) for v in candidates.values()) assert total > 50, f"Too few events generated: {total}"