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