"""Unit tests for price_gen.py.""" import datetime as dt import math import numpy as np import pytest from libs.backtest.scenarios.price_gen import ( PriceRegime, _bars_from_log_returns, _generate_regime_log_returns, generate_market_etf_paths, generate_price_paths, ) _DATES = [dt.date(2024, 1, 2) + dt.timedelta(days=i) for i in range(300)] _TRADING_DATES = [d for d in _DATES if d.weekday() < 5][:252] _RNG = np.random.default_rng(42) @pytest.mark.unit class TestPriceRegime: def test_default_jump_params(self): r = PriceRegime(annualized_drift=0.10, annualized_vol=0.15, duration_days=100) assert r.jump_prob == 0.0 assert r.jump_mean == 0.0 assert r.jump_std == 0.02 @pytest.mark.unit class TestGenerateRegimeLogReturns: def test_length_matches_n_days(self): regimes = [PriceRegime(0.10, 0.15, 100)] rets = _generate_regime_log_returns(regimes, 252, np.random.default_rng(1)) assert len(rets) == 252 def test_extends_when_regime_too_short(self): regimes = [PriceRegime(0.10, 0.15, 50)] rets = _generate_regime_log_returns(regimes, 200, np.random.default_rng(1)) assert len(rets) == 200 def test_multi_regime_transitions(self): regimes = [ PriceRegime(0.20, 0.15, 100), PriceRegime(-0.30, 0.35, 100), ] rets = _generate_regime_log_returns(regimes, 200, np.random.default_rng(2)) assert len(rets) == 200 # First half should have higher mean drift than second first_half_mean = sum(rets[:100]) / 100 second_half_mean = sum(rets[100:]) / 100 assert first_half_mean > second_half_mean # stochastic but should hold with seed def test_bear_market_negative_drift(self): regimes = [PriceRegime(-0.40, 0.30, 252)] rng = np.random.default_rng(100) rets = _generate_regime_log_returns(regimes, 252, rng) # Over 252 days, cumulative return should likely be negative cum_return = math.exp(sum(rets)) - 1 assert cum_return < 0.0 # with this seed this should hold @pytest.mark.unit class TestBarsFromLogReturns: def test_ohlcv_consistency(self): rng = np.random.default_rng(42) dates = _TRADING_DATES[:50] rets = [rng.normal(0.0004, 0.01) for _ in dates] bars = _bars_from_log_returns(rets, dates, 100.0, 0.005, 1_000_000, rng) for date, bar in bars.items(): o, h, lo, c = bar["open"], bar["high"], bar["low"], bar["close"] assert h >= max(o, c), f"high {h} < max(open {o}, close {c}) on {date}" assert lo <= min(o, c), f"low {lo} > min(open {o}, close {c}) on {date}" assert lo > 0, f"low {lo} <= 0 on {date}" assert bar["volume"] > 0 def test_returns_correct_dates(self): rng = np.random.default_rng(1) dates = _TRADING_DATES[:10] rets = [0.001] * 10 bars = _bars_from_log_returns(rets, dates, 100.0, 0.005, 1_000_000, rng) assert set(bars.keys()) == set(dates) def test_price_floor_prevents_zero(self): rng = np.random.default_rng(7) dates = _TRADING_DATES[:20] # Extreme negative returns rets = [-1.5] * 20 bars = _bars_from_log_returns(rets, dates, 100.0, 0.005, 1_000_000, rng) for bar in bars.values(): assert bar["close"] >= 0.01 assert bar["low"] >= 0.01 @pytest.mark.unit class TestGeneratePricePaths: def test_returns_correct_symbols(self): rng = np.random.default_rng(42) regimes = [PriceRegime(0.10, 0.15, 252)] tickers = ["AAPL", "MSFT", "GOOG"] result = generate_price_paths(3, None, regimes, _TRADING_DATES, rng=rng, tickers=tickers) assert isinstance(result, dict) assert set(result.keys()) == {"AAPL", "MSFT", "GOOG"} def test_auto_tickers(self): rng = np.random.default_rng(42) regimes = [PriceRegime(0.10, 0.15, 252)] result = generate_price_paths(5, None, regimes, _TRADING_DATES, rng=rng) assert len(result) == 5 for ticker in result: assert ticker.startswith("SYM") def test_all_dates_present(self): rng = np.random.default_rng(42) regimes = [PriceRegime(0.10, 0.15, 252)] result = generate_price_paths(3, [50.0, 100.0, 200.0], regimes, _TRADING_DATES[:30], rng=rng) for ticker, bars in result.items(): assert set(bars.keys()) == set(_TRADING_DATES[:30]) def test_returns_market_log_rets_when_requested(self): rng = np.random.default_rng(42) regimes = [PriceRegime(0.10, 0.15, 252)] result = generate_price_paths( 3, None, regimes, _TRADING_DATES, rng=rng, return_market_log_rets=True, ) assert isinstance(result, tuple) bars, log_rets = result assert len(log_rets) == len(_TRADING_DATES) def test_ohlcv_consistency(self): rng = np.random.default_rng(42) regimes = [PriceRegime(0.10, 0.15, 252)] result = generate_price_paths(5, None, regimes, _TRADING_DATES[:50], rng=rng) for ticker, bars in result.items(): for date, bar in bars.items(): assert bar["high"] >= max(bar["open"], bar["close"]) assert bar["low"] <= min(bar["open"], bar["close"]) assert bar["low"] > 0 assert bar["volume"] > 0 @pytest.mark.unit class TestGenerateMarketEtfPaths: def test_returns_spy_qqq_log_rets(self): rng = np.random.default_rng(42) regimes = [PriceRegime(0.12, 0.14, 252)] spy, qqq, log_rets = generate_market_etf_paths(regimes, _TRADING_DATES, rng) assert set(spy.keys()) == set(_TRADING_DATES) assert set(qqq.keys()) == set(_TRADING_DATES) assert len(log_rets) == len(_TRADING_DATES) def test_spy_and_qqq_correlated(self): """SPY and QQQ should be positively correlated in bull market.""" rng = np.random.default_rng(100) regimes = [PriceRegime(0.20, 0.14, 252)] spy, qqq, _ = generate_market_etf_paths(regimes, _TRADING_DATES, rng) spy_closes = [spy[d]["close"] for d in _TRADING_DATES] qqq_closes = [qqq[d]["close"] for d in _TRADING_DATES] spy_final_return = spy_closes[-1] / spy_closes[0] - 1 qqq_final_return = qqq_closes[-1] / qqq_closes[0] - 1 # Both should be positive in a bull market assert spy_final_return > 0 assert qqq_final_return > 0