DockerSpace / tests /test_backtest_service.py
DennisChan0909's picture
feat: add institutional signals and HF-friendly deploy
20fef51
Raw
History Blame Contribute Delete
7.37 kB
"""
Unit tests for services/backtest_service.py
Covers:
#13 — HOLD signal net_return should be 0 (no trade)
#14 — SELL signal PnL uses -actual_return (short position)
#21 — backtest logic lives in service layer, stats are correct
"""
import pytest
import inspect
from services.backtest_service import (
compute_net_return,
_build_stats,
ROUND_TRIP_COST_PCT,
BacktestError,
walk_forward_backtest,
)
# ---------------------------------------------------------------------------
# #13: HOLD net_return = 0
# ---------------------------------------------------------------------------
class TestHoldNetReturn:
"""#13 — HOLD means no trade, net_return must be 0."""
def test_hold_positive_market(self):
"""HOLD when market goes up: still 0."""
assert compute_net_return("HOLD", 5.0, ROUND_TRIP_COST_PCT) == 0.0
def test_hold_negative_market(self):
"""HOLD when market goes down: still 0."""
assert compute_net_return("HOLD", -3.0, ROUND_TRIP_COST_PCT) == 0.0
def test_hold_zero_market(self):
"""HOLD when market is flat: still 0."""
assert compute_net_return("HOLD", 0.0, ROUND_TRIP_COST_PCT) == 0.0
# ---------------------------------------------------------------------------
# #14: SELL net_return = -actual_return - cost (short position)
# ---------------------------------------------------------------------------
class TestSellNetReturn:
"""#14 — SELL = short position, profits when price drops."""
def test_sell_price_drops(self):
"""SELL when price drops 3%: profit = +3% - cost."""
actual = -3.0
result = compute_net_return("SELL", actual, ROUND_TRIP_COST_PCT)
expected = -actual - ROUND_TRIP_COST_PCT # 3.0 - 0.471 = 2.529
assert abs(result - expected) < 1e-10
def test_sell_price_rises(self):
"""SELL when price rises 2%: loss = -2% - cost."""
actual = 2.0
result = compute_net_return("SELL", actual, ROUND_TRIP_COST_PCT)
expected = -actual - ROUND_TRIP_COST_PCT # -2.0 - 0.471 = -2.471
assert abs(result - expected) < 1e-10
def test_sell_price_flat(self):
"""SELL when price flat: loss = -cost only."""
result = compute_net_return("SELL", 0.0, ROUND_TRIP_COST_PCT)
expected = -ROUND_TRIP_COST_PCT
assert abs(result - expected) < 1e-10
# ---------------------------------------------------------------------------
# BUY net_return (unchanged behaviour, regression test)
# ---------------------------------------------------------------------------
class TestBuyNetReturn:
"""BUY = long position, profits when price rises."""
def test_buy_price_rises(self):
actual = 5.0
result = compute_net_return("BUY", actual, ROUND_TRIP_COST_PCT)
expected = actual - ROUND_TRIP_COST_PCT
assert abs(result - expected) < 1e-10
def test_buy_price_drops(self):
actual = -2.0
result = compute_net_return("BUY", actual, ROUND_TRIP_COST_PCT)
expected = actual - ROUND_TRIP_COST_PCT
assert abs(result - expected) < 1e-10
# ---------------------------------------------------------------------------
# #21: _build_stats aggregation
# ---------------------------------------------------------------------------
class TestBuildStats:
"""#21 — Stats aggregation logic (extracted to service layer)."""
@pytest.fixture
def sample_trades(self):
"""Mix of BUY / SELL / HOLD trades for stat verification."""
return [
{
"date": "2025-01-02",
"signal": "BUY",
"buy_prob": 0.75,
"close": 100.0,
"future_close": 105.0,
"actual_return_pct": 5.0,
"net_return_pct": round(5.0 - ROUND_TRIP_COST_PCT, 2),
"correct": True,
},
{
"date": "2025-01-03",
"signal": "SELL",
"buy_prob": 0.20,
"close": 100.0,
"future_close": 97.0,
"actual_return_pct": -3.0,
"net_return_pct": round(3.0 - ROUND_TRIP_COST_PCT, 2),
"correct": True,
},
{
"date": "2025-01-04",
"signal": "HOLD",
"buy_prob": 0.50,
"close": 100.0,
"future_close": 102.0,
"actual_return_pct": 2.0,
"net_return_pct": 0.0,
"correct": False, # HOLD 不算 correct/incorrect in non_holds
},
{
"date": "2025-01-05",
"signal": "SELL",
"buy_prob": 0.30,
"close": 100.0,
"future_close": 103.0,
"actual_return_pct": 3.0,
"net_return_pct": round(-3.0 - ROUND_TRIP_COST_PCT, 2),
"correct": False, # predicted down, went up
},
]
def test_signal_counts(self, sample_trades):
result = _build_stats("TEST", 6, sample_trades)
stats = result["stats"]
assert stats["buy_signals"] == 1
assert stats["sell_signals"] == 2
assert stats["hold_signals"] == 1
assert stats["total_days"] == 4
def test_hold_not_in_total_return(self, sample_trades):
"""HOLD trades contribute 0 to total_return_if_followed."""
result = _build_stats("TEST", 6, sample_trades)
non_hold_returns = sum(
t["net_return_pct"] for t in sample_trades if t["signal"] != "HOLD"
)
assert result["stats"]["total_return_if_followed"] == round(non_hold_returns, 2)
def test_signal_accuracy_excludes_hold(self, sample_trades):
"""Signal accuracy only considers BUY + SELL, not HOLD."""
result = _build_stats("TEST", 6, sample_trades)
# 3 non-hold trades: 2 correct (BUY correct, SELL#1 correct), 1 wrong (SELL#2)
expected = round(2 / 3 * 100, 1)
assert result["stats"]["signal_accuracy_overall"] == expected
def test_empty_results(self):
"""Empty results should not crash."""
result = _build_stats("TEST", 6, [])
assert result["stats"]["total_days"] == 0
assert result["stats"]["best_trade_pct"] == 0
assert result["stats"]["worst_trade_pct"] == 0
def test_metadata_fields(self, sample_trades):
result = _build_stats("TEST", 6, sample_trades)
assert result["stock_no"] == "TEST"
assert result["months"] == 6
assert result["walk_forward"] is True
assert "trading_cost_pct" in result
assert "cost_breakdown" in result
# ---------------------------------------------------------------------------
# BacktestError (used by router to raise 404)
# ---------------------------------------------------------------------------
class TestBacktestError:
def test_is_exception(self):
with pytest.raises(BacktestError):
raise BacktestError("No data")
def test_message(self):
err = BacktestError("No data found for 9999")
assert "9999" in str(err)
class TestBacktestLeakageGuards:
def test_backtest_does_not_use_news_overlay(self):
source = inspect.getsource(walk_forward_backtest).lower()
assert "news" not in source