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.
198 lines
7.0 KiB
Python
198 lines
7.0 KiB
Python
"""Unit tests for labeler module."""
|
|
from __future__ import annotations
|
|
|
|
import datetime as dt
|
|
from decimal import Decimal
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
import pytest
|
|
|
|
from libs.labeler.label_generator import _compute_labels_from_bars
|
|
from libs.labeler.reaction_date import compute_reaction_date
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestComputeReactionDate:
|
|
"""Tests for reaction date calculation."""
|
|
|
|
def test_pre_market_on_trading_day_returns_same_day(self) -> None:
|
|
"""Filing before market open on a trading day → same day reaction."""
|
|
# 2026-01-02 is a Friday (trading day)
|
|
d = dt.date(2026, 1, 2)
|
|
result = compute_reaction_date(d, "pre_market")
|
|
assert result == d
|
|
|
|
def test_regular_hours_on_trading_day_returns_same_day(self) -> None:
|
|
"""Filing during regular hours on a trading day → same day reaction."""
|
|
d = dt.date(2026, 1, 2)
|
|
result = compute_reaction_date(d, "regular_hours")
|
|
assert result == d
|
|
|
|
def test_post_market_returns_next_trading_day(self) -> None:
|
|
"""Filing after market close → next trading day reaction."""
|
|
d = dt.date(2026, 1, 2)
|
|
result = compute_reaction_date(d, "post_market")
|
|
assert result > d
|
|
|
|
def test_unknown_returns_next_trading_day(self) -> None:
|
|
"""Unknown bucket → conservative: next trading day."""
|
|
d = dt.date(2026, 1, 2)
|
|
result = compute_reaction_date(d, "unknown")
|
|
assert result > d
|
|
|
|
def test_pre_market_on_weekend_returns_next_trading_day(self) -> None:
|
|
"""Pre-market filing on weekend (non-trading day) → next trading day."""
|
|
saturday = dt.date(2026, 1, 3) # Saturday
|
|
result = compute_reaction_date(saturday, "pre_market")
|
|
assert result > saturday
|
|
|
|
def test_post_market_on_friday_returns_monday(self) -> None:
|
|
"""Post-market on Friday → next Monday (assuming no holiday)."""
|
|
friday = dt.date(2026, 1, 2) # 2026-01-02 is a Friday
|
|
result = compute_reaction_date(friday, "post_market")
|
|
# Next trading day after Friday is Monday
|
|
assert result.weekday() == 0 # Monday
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestComputeLabelsFromBars:
|
|
"""Tests for forward-return label computation."""
|
|
|
|
def _bar(self, close: float, high: float | None = None, low: float | None = None) -> dict:
|
|
return {
|
|
"open": close * 0.99,
|
|
"high": high if high is not None else close * 1.02,
|
|
"low": low if low is not None else close * 0.98,
|
|
"close": close,
|
|
}
|
|
|
|
def test_1d_return_calculation(self) -> None:
|
|
entry = Decimal("100")
|
|
bars = [self._bar(105)] # +5%
|
|
result = _compute_labels_from_bars(entry, bars, 1)
|
|
assert abs(float(result["fwd_return"]) - 0.05) < 0.001
|
|
|
|
def test_mfe_is_max_high_minus_entry(self) -> None:
|
|
entry = Decimal("100")
|
|
bars = [
|
|
self._bar(101, high=105),
|
|
self._bar(103, high=108),
|
|
self._bar(102, high=104),
|
|
]
|
|
result = _compute_labels_from_bars(entry, bars, 3)
|
|
# Max high = 108, so MFE = (108-100)/100 = 0.08
|
|
assert abs(float(result["mfe"]) - 0.08) < 0.001
|
|
|
|
def test_mae_is_min_low_minus_entry(self) -> None:
|
|
entry = Decimal("100")
|
|
bars = [
|
|
self._bar(99, low=97),
|
|
self._bar(98, low=95),
|
|
self._bar(100, low=98),
|
|
]
|
|
result = _compute_labels_from_bars(entry, bars, 3)
|
|
# Min low = 95, so MAE = (95-100)/100 = -0.05
|
|
assert abs(float(result["mae"]) - (-0.05)) < 0.001
|
|
|
|
def test_hit_pos_1r_true_when_high_exceeds_threshold(self) -> None:
|
|
entry = Decimal("100")
|
|
bars = [self._bar(99, high=101.5)] # +1.5% > 1R threshold
|
|
result = _compute_labels_from_bars(entry, bars, 1)
|
|
assert result["hit_pos_1r"] is True
|
|
|
|
def test_hit_pos_1r_false_when_high_below_threshold(self) -> None:
|
|
entry = Decimal("100")
|
|
bars = [self._bar(99, high=100.5)] # +0.5% < 1R threshold
|
|
result = _compute_labels_from_bars(entry, bars, 1)
|
|
assert result["hit_pos_1r"] is False
|
|
|
|
def test_close_up_after_3d_true_when_final_close_above_entry(self) -> None:
|
|
entry = Decimal("100")
|
|
bars = [self._bar(98), self._bar(101), self._bar(103)]
|
|
result = _compute_labels_from_bars(entry, bars, 3)
|
|
assert result["close_up"] is True
|
|
|
|
def test_empty_bars_returns_empty_dict(self) -> None:
|
|
result = _compute_labels_from_bars(Decimal("100"), [], 3)
|
|
assert result == {}
|
|
|
|
|
|
@pytest.mark.unit
|
|
class TestGenerateLabels:
|
|
"""Tests for async generate_labels function."""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_generate_labels_with_valid_prices(self) -> None:
|
|
"""generate_labels returns EventLabel with ok status when prices available."""
|
|
from libs.labeler.label_generator import generate_labels
|
|
|
|
# Mock event
|
|
mock_event = MagicMock()
|
|
mock_event.event_id = "EVT::test::001"
|
|
mock_event.event_date = dt.date(2026, 1, 5) # Monday
|
|
mock_event.filing_time_bucket = "post_market"
|
|
|
|
# Mock price service
|
|
mock_price_svc = AsyncMock()
|
|
mock_bar = MagicMock()
|
|
mock_bar.model_dump.return_value = {
|
|
"date": "2026-01-07", # Wednesday = entry_date
|
|
"open": 100.0,
|
|
"high": 105.0,
|
|
"low": 98.0,
|
|
"close": 103.0,
|
|
}
|
|
# Create 8 bars for look-ahead
|
|
bars = []
|
|
for i in range(8):
|
|
b = MagicMock()
|
|
date = dt.date(2026, 1, 7) + dt.timedelta(days=i)
|
|
b.model_dump.return_value = {
|
|
"date": date.isoformat(),
|
|
"open": 100.0 + i,
|
|
"high": 105.0 + i,
|
|
"low": 98.0,
|
|
"close": 103.0 + i,
|
|
}
|
|
bars.append(b)
|
|
|
|
mock_resp = MagicMock()
|
|
mock_resp.bars = bars
|
|
mock_price_svc.get_daily_bars = AsyncMock(return_value=mock_resp)
|
|
|
|
mock_session = AsyncMock()
|
|
|
|
label = await generate_labels(
|
|
session=mock_session,
|
|
event=mock_event,
|
|
price_svc=mock_price_svc,
|
|
ticker="AAPL",
|
|
)
|
|
|
|
assert label.event_id == "EVT::test::001"
|
|
assert label.label_status in ("ok", "truncated")
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_generate_labels_unavailable_when_no_price_data(self) -> None:
|
|
"""generate_labels returns 'unavailable' status on price fetch error."""
|
|
from libs.labeler.label_generator import generate_labels
|
|
|
|
mock_event = MagicMock()
|
|
mock_event.event_id = "EVT::test::002"
|
|
mock_event.event_date = dt.date(2026, 1, 5)
|
|
mock_event.filing_time_bucket = "post_market"
|
|
|
|
mock_price_svc = AsyncMock()
|
|
mock_price_svc.get_daily_bars = AsyncMock(side_effect=Exception("Oracle unavailable"))
|
|
|
|
mock_session = AsyncMock()
|
|
|
|
label = await generate_labels(
|
|
session=mock_session,
|
|
event=mock_event,
|
|
price_svc=mock_price_svc,
|
|
ticker="AAPL",
|
|
)
|
|
|
|
assert label.label_status == "unavailable"
|