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.
93 lines
2.9 KiB
Python
93 lines
2.9 KiB
Python
"""Integration test: event → label generation pipeline (real Oracle + DB)."""
|
|
from __future__ import annotations
|
|
|
|
import datetime as dt
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
import pytest
|
|
|
|
pytestmark = pytest.mark.integration
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_event() -> MagicMock:
|
|
event = MagicMock()
|
|
event.event_id = "EVT::sec::ISSUER::0000320193::2026-01-29::earnings_release::0"
|
|
event.event_date = dt.date(2026, 1, 29)
|
|
event.filed_at_utc = dt.datetime(2026, 1, 29, 22, 0, tzinfo=dt.UTC) # post_market
|
|
event.symbol_id = "SYM::AAPL::NASDAQ"
|
|
event.status = "valid"
|
|
return event
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_price_bars() -> list[MagicMock]:
|
|
bars = []
|
|
for i in range(8):
|
|
bar = MagicMock()
|
|
date = dt.date(2026, 1, 30) + dt.timedelta(days=i)
|
|
bar.model_dump.return_value = {
|
|
"date": date.isoformat(),
|
|
"open": 220.0 + i * 0.5,
|
|
"high": 225.0 + i * 0.5,
|
|
"low": 218.0,
|
|
"close": 222.0 + i * 0.5,
|
|
"volume": 1000000,
|
|
}
|
|
bars.append(bar)
|
|
return bars
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_label_pipeline_end_to_end(mock_event: MagicMock, mock_price_bars: list) -> None:
|
|
"""Test full label generation: event → reaction_date → entry → labels."""
|
|
from libs.labeler.label_generator import LABEL_VERSION, generate_labels
|
|
from libs.labeler.reaction_date import compute_reaction_date
|
|
|
|
# Verify reaction_date logic
|
|
reaction_date = compute_reaction_date(mock_event.event_date, mock_event.filing_time_bucket)
|
|
assert reaction_date > mock_event.event_date # post_market → next day
|
|
|
|
# Setup mock price service
|
|
mock_price_resp = MagicMock()
|
|
mock_price_resp.bars = mock_price_bars
|
|
mock_price_svc = AsyncMock()
|
|
mock_price_svc.get_daily_bars = AsyncMock(return_value=mock_price_resp)
|
|
mock_session = AsyncMock()
|
|
|
|
label = await generate_labels(
|
|
session=mock_session,
|
|
event=mock_event,
|
|
price_svc=mock_price_svc,
|
|
ticker="AAPL",
|
|
)
|
|
|
|
assert label is not None
|
|
assert label.event_id == mock_event.event_id
|
|
assert label.label_version == LABEL_VERSION
|
|
assert label.label_status in ("ok", "truncated")
|
|
assert label.reaction_date == reaction_date
|
|
assert label.entry_price is not None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_label_pipeline_handles_missing_price_data(mock_event: MagicMock) -> None:
|
|
"""Label pipeline handles Oracle price unavailability gracefully."""
|
|
from libs.labeler.label_generator import generate_labels
|
|
|
|
mock_price_svc = AsyncMock()
|
|
mock_price_svc.get_daily_bars = AsyncMock(
|
|
side_effect=Exception("Oracle connection refused")
|
|
)
|
|
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.event_id == mock_event.event_id
|