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

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