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.

327 lines
12 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.filed_at_utc = dt.datetime(2026, 1, 5, 22, 0, tzinfo=dt.UTC) # 5PM ET = 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
for an event with a past entry_date (Oracle has the data; 404 is real).
"""
from libs.labeler.label_generator import generate_labels
mock_event = MagicMock()
mock_event.event_id = "EVT::test::002"
# Far in the past — entry_date will also be in the past
mock_event.event_date = dt.date(2020, 1, 5)
mock_event.filed_at_utc = dt.datetime(2020, 1, 5, 22, 0, tzinfo=dt.UTC)
mock_price_svc = AsyncMock()
mock_price_svc.get_daily_bars = AsyncMock(
side_effect=Exception("Not found: /api/v1/price/data/AAPL")
)
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"
assert label.entry_date is None
@pytest.mark.asyncio
async def test_generate_labels_pending_on_future_entry_date_404(self) -> None:
"""Reproduces the 2026-05-08 RKLB/SNDK/AKAM bug: pre-market label_generator
run requests prices for entry_date in today/future window; Oracle 404s
because those bars don't exist yet. Must return 'pending' (not 'unavailable')
with entry_date preserved so the post-close run can regenerate it AND the
live snapshot includes the candidate.
"""
from libs.labeler.label_generator import generate_labels
# Post-market filing dated TODAY → reaction_date = next trading day
# entry_date for next_open_after_reaction_close = day after that → future
today = dt.date.today()
mock_event = MagicMock()
mock_event.event_id = "EVT::test::future"
mock_event.event_date = today
# 22:00 UTC = post_market in ET
mock_event.filed_at_utc = dt.datetime.combine(
today, dt.time(22, 0), tzinfo=dt.UTC
)
mock_price_svc = AsyncMock()
mock_price_svc.get_daily_bars = AsyncMock(
side_effect=Exception("Not found: /api/v1/price/data/RKLB")
)
mock_session = AsyncMock()
label = await generate_labels(
session=mock_session,
event=mock_event,
price_svc=mock_price_svc,
ticker="RKLB",
)
assert label.label_status == "pending", (
"Future-window 404 must yield 'pending' so the post-close pipeline "
"regenerates it and the live snapshot includes today's candidates"
)
# entry_date must be preserved (downstream filters need it; old code set it None)
assert label.entry_date is not None
assert label.entry_date >= today
assert label.reaction_date is not None
@pytest.mark.asyncio
async def test_generate_labels_pre_market_same_day_reaction(self) -> None:
"""Pre-market filing → reaction_date == event_date."""
from libs.labeler.label_generator import generate_labels
mock_event = MagicMock()
mock_event.event_id = "EVT::test::003"
mock_event.event_date = dt.date(2026, 1, 5) # Monday
mock_event.filed_at_utc = dt.datetime(2026, 1, 5, 12, 0, tzinfo=dt.UTC) # 7AM ET = pre_market
# Create bars starting from entry_date (next day after reaction=event_date)
bars = []
for i in range(8):
b = MagicMock()
date = dt.date(2026, 1, 6) + 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 = AsyncMock()
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",
)
# Pre-market on Monday → reaction_date == event_date (Monday)
assert label.reaction_date == dt.date(2026, 1, 5)
@pytest.mark.asyncio
async def test_generate_labels_filed_at_utc_none_defaults_post_market(self) -> None:
"""filed_at_utc=None → 'unknown' → next trading day."""
from libs.labeler.label_generator import generate_labels
mock_event = MagicMock()
mock_event.event_id = "EVT::test::004"
mock_event.event_date = dt.date(2026, 1, 5) # Monday
mock_event.filed_at_utc = None
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 = AsyncMock()
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",
)
# None → unknown → next trading day after Monday = Tuesday
assert label.reaction_date > dt.date(2026, 1, 5)