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.

168 lines
6.4 KiB
Python

"""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