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.
270 lines
10 KiB
Python
270 lines
10 KiB
Python
"""Unit tests for libs/backtest/selector.py."""
|
|
from __future__ import annotations
|
|
|
|
import datetime as dt
|
|
from zoneinfo import ZoneInfo
|
|
|
|
import pytest
|
|
|
|
from libs.backtest.domain import EventTypeProfile, SignalConfig, UniverseConfig
|
|
|
|
_UTC = ZoneInfo("UTC")
|
|
|
|
|
|
def _make_raw_row(**kwargs) -> dict:
|
|
defaults = {
|
|
"event_id": "EVT::TEST::001",
|
|
"symbol": "AAPL",
|
|
"issuer_id": "ISSUER::0000320193",
|
|
"score": 0.75,
|
|
"sector": "Technology",
|
|
"event_type": "earnings",
|
|
"event_timestamp": "2026-01-05T21:00:00+00:00",
|
|
"filing_time_bucket": "post_market",
|
|
"reaction_date": "2026-01-06",
|
|
"entry_date": "2026-01-07",
|
|
"entry_price": 150.0,
|
|
"avg_dollar_volume": 5_000_000.0,
|
|
"atr_14": 3.5,
|
|
}
|
|
defaults.update(kwargs)
|
|
return defaults
|
|
|
|
|
|
class TestBuildCandidate:
|
|
def test_basic(self):
|
|
from libs.backtest.selector import build_candidate
|
|
|
|
row = _make_raw_row()
|
|
c = build_candidate(row)
|
|
assert c is not None
|
|
assert c.symbol == "AAPL"
|
|
assert c.score == 0.75
|
|
assert c.execution_date == dt.date(2026, 1, 7)
|
|
assert c.event_timestamp.tzinfo is not None
|
|
|
|
def test_null_timestamp_returns_none(self):
|
|
from libs.backtest.selector import build_candidate
|
|
|
|
row = _make_raw_row(event_timestamp=None)
|
|
assert build_candidate(row) is None
|
|
|
|
def test_zero_entry_price_returns_none(self):
|
|
from libs.backtest.selector import build_candidate
|
|
|
|
row = _make_raw_row(entry_price=0.0)
|
|
assert build_candidate(row) is None
|
|
|
|
def test_missing_entry_price_returns_none(self):
|
|
from libs.backtest.selector import build_candidate
|
|
|
|
row = _make_raw_row()
|
|
del row["entry_price"]
|
|
assert build_candidate(row) is None
|
|
|
|
def test_null_exec_date_returns_none(self):
|
|
from libs.backtest.selector import build_candidate
|
|
|
|
row = _make_raw_row()
|
|
del row["entry_date"]
|
|
assert build_candidate(row) is None
|
|
|
|
def test_sector_defaults_to_unknown(self):
|
|
from libs.backtest.selector import build_candidate
|
|
|
|
row = _make_raw_row(sector=None)
|
|
c = build_candidate(row)
|
|
assert c is not None
|
|
assert c.sector == "UNKNOWN"
|
|
|
|
def test_score_bucket_classification(self):
|
|
from libs.backtest.selector import build_candidate
|
|
|
|
c = build_candidate(_make_raw_row(score=0.85))
|
|
assert c.score_bucket == "high"
|
|
|
|
c = build_candidate(_make_raw_row(score=0.65))
|
|
assert c.score_bucket == "medium_high"
|
|
|
|
c = build_candidate(_make_raw_row(score=0.45))
|
|
assert c.score_bucket == "medium"
|
|
|
|
c = build_candidate(_make_raw_row(score=0.25))
|
|
assert c.score_bucket == "medium_low"
|
|
|
|
c = build_candidate(_make_raw_row(score=0.10))
|
|
assert c.score_bucket == "low"
|
|
|
|
|
|
class TestRankCandidates:
|
|
def test_sorted_by_score_desc(self):
|
|
from libs.backtest.selector import build_candidate, rank_candidates
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="A", score=0.5, avg_dollar_volume=1e6),
|
|
_make_raw_row(symbol="B", score=0.8, avg_dollar_volume=1e6),
|
|
_make_raw_row(symbol="C", score=0.6, avg_dollar_volume=1e6),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows]
|
|
ranked = rank_candidates([c for c in candidates if c])
|
|
assert ranked[0].symbol == "B"
|
|
assert ranked[1].symbol == "C"
|
|
assert ranked[2].symbol == "A"
|
|
|
|
def test_tiebreak_by_avg_dollar_volume(self):
|
|
from libs.backtest.selector import build_candidate, rank_candidates
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="A", score=0.7, avg_dollar_volume=1e6),
|
|
_make_raw_row(symbol="B", score=0.7, avg_dollar_volume=5e6),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows]
|
|
ranked = rank_candidates([c for c in candidates if c])
|
|
assert ranked[0].symbol == "B" # higher avg_dollar_volume
|
|
|
|
def test_tiebreak_by_symbol_asc(self):
|
|
from libs.backtest.selector import build_candidate, rank_candidates
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="Z", score=0.7, avg_dollar_volume=1e6),
|
|
_make_raw_row(symbol="A", score=0.7, avg_dollar_volume=1e6),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows]
|
|
ranked = rank_candidates([c for c in candidates if c])
|
|
assert ranked[0].symbol == "A"
|
|
|
|
def test_deterministic(self):
|
|
from libs.backtest.selector import build_candidate, rank_candidates
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="C", score=0.9),
|
|
_make_raw_row(symbol="A", score=0.7),
|
|
_make_raw_row(symbol="B", score=0.8),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows]
|
|
r1 = rank_candidates([c for c in candidates if c])
|
|
r2 = rank_candidates([c for c in candidates if c])
|
|
assert [c.symbol for c in r1] == [c.symbol for c in r2]
|
|
|
|
|
|
class TestFilterCandidates:
|
|
def test_score_threshold(self):
|
|
from libs.backtest.selector import build_candidate, filter_by_score
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="A", score=0.3),
|
|
_make_raw_row(symbol="B", score=0.7),
|
|
_make_raw_row(symbol="C", score=0.5),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows if build_candidate(r)]
|
|
filtered = filter_by_score(candidates, score_threshold=0.5)
|
|
assert len(filtered) == 2
|
|
assert all(c.score >= 0.5 for c in filtered)
|
|
|
|
def test_min_price_filter(self):
|
|
from libs.backtest.selector import build_candidate, filter_by_universe
|
|
|
|
u = UniverseConfig(min_price=100.0, min_avg_dollar_volume=0)
|
|
rows = [
|
|
_make_raw_row(symbol="CHEAP", entry_price=50.0),
|
|
_make_raw_row(symbol="OK", entry_price=150.0),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows if build_candidate(r)]
|
|
filtered = filter_by_universe(candidates, u)
|
|
assert len(filtered) == 1
|
|
assert filtered[0].symbol == "OK"
|
|
|
|
def test_min_adv_filter(self):
|
|
from libs.backtest.selector import build_candidate, filter_by_universe
|
|
|
|
u = UniverseConfig(min_price=0, min_avg_dollar_volume=2_000_000)
|
|
rows = [
|
|
_make_raw_row(symbol="ILLIQUID", avg_dollar_volume=500_000),
|
|
_make_raw_row(symbol="LIQUID", avg_dollar_volume=5_000_000),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows if build_candidate(r)]
|
|
filtered = filter_by_universe(candidates, u)
|
|
assert len(filtered) == 1
|
|
assert filtered[0].symbol == "LIQUID"
|
|
|
|
def test_truncate(self):
|
|
from libs.backtest.selector import build_candidate, rank_candidates, truncate_candidates
|
|
|
|
rows = [_make_raw_row(symbol=s, score=0.9 - i * 0.1) for i, s in enumerate("ABCDE")]
|
|
candidates = rank_candidates([build_candidate(r) for r in rows if build_candidate(r)])
|
|
truncated = truncate_candidates(candidates, max_per_day=3)
|
|
assert len(truncated) == 3
|
|
|
|
|
|
class TestFilterByEventType:
|
|
def test_disabled_event_type_filtered(self):
|
|
from libs.backtest.selector import build_candidate, filter_by_event_type
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="A", event_type="earnings_release"),
|
|
_make_raw_row(symbol="B", event_type="management_change"),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows if build_candidate(r)]
|
|
profiles = {
|
|
"management_change": EventTypeProfile(enabled=False),
|
|
}
|
|
filtered = filter_by_event_type(candidates, profiles)
|
|
assert len(filtered) == 1
|
|
assert filtered[0].symbol == "A"
|
|
|
|
def test_per_type_score_threshold(self):
|
|
from libs.backtest.selector import build_candidate, filter_by_event_type
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="A", event_type="earnings_release", score=0.55),
|
|
_make_raw_row(symbol="B", event_type="earnings_release", score=0.75),
|
|
]
|
|
candidates = [build_candidate(r) for r in rows if build_candidate(r)]
|
|
profiles = {
|
|
"earnings_release": EventTypeProfile(score_threshold_override=0.6),
|
|
}
|
|
filtered = filter_by_event_type(candidates, profiles)
|
|
assert len(filtered) == 1
|
|
assert filtered[0].symbol == "B"
|
|
|
|
def test_no_profiles_passthrough(self):
|
|
from libs.backtest.selector import build_candidate, filter_by_event_type
|
|
|
|
rows = [_make_raw_row(symbol="A")]
|
|
candidates = [build_candidate(r) for r in rows if build_candidate(r)]
|
|
assert filter_by_event_type(candidates, {}) == candidates
|
|
|
|
|
|
class TestSelectCandidates:
|
|
def test_full_pipeline(self):
|
|
from libs.backtest.selector import select_candidates
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="A", score=0.9, avg_dollar_volume=5e6, entry_price=100.0),
|
|
_make_raw_row(symbol="B", score=0.3, avg_dollar_volume=5e6, entry_price=100.0), # below threshold
|
|
_make_raw_row(symbol="C", score=0.8, avg_dollar_volume=1e4, entry_price=100.0), # low ADV
|
|
_make_raw_row(symbol="D", score=0.7, avg_dollar_volume=5e6, entry_price=2.0), # below min_price
|
|
]
|
|
u = UniverseConfig(min_price=5.0, min_avg_dollar_volume=1_000_000)
|
|
s = SignalConfig(score_threshold=0.5, max_candidates_per_day=10)
|
|
result = select_candidates(rows, u, s)
|
|
symbols = [c.symbol for c in result]
|
|
assert "A" in symbols
|
|
assert "B" not in symbols # below threshold
|
|
assert "C" not in symbols # low ADV
|
|
assert "D" not in symbols # below min_price
|
|
|
|
def test_pipeline_with_event_type_profiles(self):
|
|
from libs.backtest.selector import select_candidates
|
|
|
|
rows = [
|
|
_make_raw_row(symbol="A", score=0.9, avg_dollar_volume=5e6, entry_price=100.0, event_type="earnings_release"),
|
|
_make_raw_row(symbol="B", score=0.7, avg_dollar_volume=5e6, entry_price=100.0, event_type="management_change"),
|
|
]
|
|
u = UniverseConfig(min_price=5.0, min_avg_dollar_volume=1_000_000)
|
|
s = SignalConfig(score_threshold=0.5, max_candidates_per_day=10)
|
|
profiles = {"management_change": EventTypeProfile(enabled=False)}
|
|
result = select_candidates(rows, u, s, event_type_profiles=profiles)
|
|
assert len(result) == 1
|
|
assert result[0].symbol == "A"
|