Cognitive-rag / 01_data_ingestion / tests / test_events.py
test_events.py
Raw
"""Tests for the SSE event system."""

from __future__ import annotations

import asyncio
import json
import sys
from pathlib import Path

import pytest

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from ingestion.events import (
    CallbackEmitter,
    FileCompleted,
    FileFailed,
    FileProgress,
    FileSkipped,
    FileStarted,
    FileEmpty,
    IngestionCompleted,
    IngestionProgress,
    IngestionStarted,
    LogEmitter,
    MultiEmitter,
    NullEmitter,
    QueueEmitter,
    event_to_dict,
    event_to_sse,
)


# ---------------------------------------------------------------------------
# Event dataclass tests
# ---------------------------------------------------------------------------

class TestEventDataclasses:
    def test_ingestion_started_defaults(self):
        e = IngestionStarted(total_files=10, skipped_files=[{"path": "a.exe", "reason": "unsupported"}])
        assert e.event_type == "ingestion_started"
        assert e.total_files == 10
        assert len(e.skipped_files) == 1

    def test_file_started_fields(self):
        e = FileStarted(file_path="docs/a.pdf", file_index=0, total_files=5, modality="text", filetype="pdf")
        assert e.event_type == "file_started"
        assert e.file_path == "docs/a.pdf"

    def test_file_progress_optional_percent(self):
        e1 = FileProgress(file_path="video.mp4", detail="Chunk 2/5")
        assert e1.percent is None
        e2 = FileProgress(file_path="video.mp4", detail="Chunk 2/5", percent=40.0)
        assert e2.percent == 40.0

    def test_file_completed(self):
        e = FileCompleted(file_path="a.txt", records_emitted=3, duration_sec=1.5)
        assert e.event_type == "file_completed"
        assert e.records_emitted == 3

    def test_file_failed(self):
        e = FileFailed(file_path="bad.pdf", error="corrupt", modality="text")
        assert e.event_type == "file_failed"

    def test_file_skipped(self):
        e = FileSkipped(file_path="a.exe", reason="Unsupported file type: .exe")
        assert e.event_type == "file_skipped"

    def test_file_empty(self):
        e = FileEmpty(file_path="blank.png", modality="image", reason="No text found")
        assert e.event_type == "file_empty"

    def test_ingestion_progress(self):
        e = IngestionProgress(files_completed=3, total_files=10, records_emitted=50, elapsed_sec=12.5, eta_sec=30.0)
        assert e.event_type == "ingestion_progress"
        assert e.eta_sec == 30.0

    def test_ingestion_completed(self):
        e = IngestionCompleted(stats={"counts": {"records_emitted": 100}})
        assert e.event_type == "ingestion_completed"
        assert e.stats["counts"]["records_emitted"] == 100


# ---------------------------------------------------------------------------
# Serialization tests
# ---------------------------------------------------------------------------

class TestSerialization:
    def test_event_to_dict(self):
        e = FileStarted(file_path="a.txt", file_index=0, total_files=1, modality="text", filetype="txt")
        d = event_to_dict(e)
        assert d["event_type"] == "file_started"
        assert d["file_path"] == "a.txt"
        assert isinstance(d, dict)

    def test_event_to_sse_format(self):
        e = FileCompleted(file_path="a.txt", records_emitted=1, duration_sec=0.5)
        sse = event_to_sse(e)
        assert sse.startswith("event: file_completed\n")
        assert "data: " in sse
        assert sse.endswith("\n\n")
        # Data should be valid JSON
        data_line = sse.split("data: ")[1].rstrip("\n")
        parsed = json.loads(data_line)
        assert parsed["records_emitted"] == 1

    def test_all_event_types_serializable(self):
        events = [
            IngestionStarted(total_files=1),
            FileStarted(file_path="a", file_index=0, total_files=1, modality="text", filetype="txt"),
            FileProgress(file_path="a", detail="working"),
            FileCompleted(file_path="a", records_emitted=1, duration_sec=0.1),
            FileFailed(file_path="a", error="oops", modality="text"),
            FileSkipped(file_path="a", reason="nope"),
            FileEmpty(file_path="a", modality="text", reason="empty"),
            IngestionProgress(files_completed=1, total_files=1, records_emitted=1, elapsed_sec=1.0),
            IngestionCompleted(stats={}),
        ]
        for event in events:
            sse = event_to_sse(event)
            assert f"event: {event.event_type}" in sse
            data = json.loads(sse.split("data: ")[1].rstrip("\n"))
            assert data["event_type"] == event.event_type


# ---------------------------------------------------------------------------
# Emitter tests
# ---------------------------------------------------------------------------

class TestNullEmitter:
    def test_emit_is_noop(self):
        emitter = NullEmitter()
        emitter.emit(FileStarted(file_path="a", file_index=0, total_files=1, modality="text", filetype="txt"))
        # Should not raise


class TestLogEmitter:
    def test_emit_prints(self, capsys):
        emitter = LogEmitter()
        emitter.emit(IngestionStarted(total_files=5))
        captured = capsys.readouterr()
        assert "5 files" in captured.out

    def test_all_event_types(self, capsys):
        emitter = LogEmitter()
        events = [
            IngestionStarted(total_files=1, skipped_files=[{"path": "a.exe", "reason": "unsupported"}]),
            FileStarted(file_path="a.txt", file_index=0, total_files=1, modality="text", filetype="txt"),
            FileProgress(file_path="a.txt", detail="working", percent=50.0),
            FileCompleted(file_path="a.txt", records_emitted=1, duration_sec=0.1),
            FileFailed(file_path="b.pdf", error="corrupt", modality="text"),
            FileSkipped(file_path="c.exe", reason="unsupported"),
            FileEmpty(file_path="d.png", modality="image", reason="no text"),
            IngestionProgress(files_completed=1, total_files=1, records_emitted=1, elapsed_sec=1.0, eta_sec=0.0),
            IngestionCompleted(stats={"counts": {"records_emitted": 1, "files_succeeded": 1}}),
        ]
        for event in events:
            emitter.emit(event)
        captured = capsys.readouterr()
        assert "[ingest]" in captured.out


class TestCallbackEmitter:
    def test_calls_callback(self):
        received = []
        emitter = CallbackEmitter(received.append)
        event = FileStarted(file_path="a", file_index=0, total_files=1, modality="text", filetype="txt")
        emitter.emit(event)
        assert len(received) == 1
        assert received[0] is event


class TestMultiEmitter:
    def test_broadcasts_to_all(self):
        r1, r2 = [], []
        e1 = CallbackEmitter(r1.append)
        e2 = CallbackEmitter(r2.append)
        multi = MultiEmitter(e1, e2)
        event = FileCompleted(file_path="a", records_emitted=1, duration_sec=0.1)
        multi.emit(event)
        assert len(r1) == 1
        assert len(r2) == 1


# ---------------------------------------------------------------------------
# QueueEmitter tests
# ---------------------------------------------------------------------------

class TestQueueEmitter:
    def test_emit_and_stream(self):
        emitter = QueueEmitter()
        event = FileCompleted(file_path="a.txt", records_emitted=1, duration_sec=0.1)
        emitter.emit(event)
        emitter.close()

        async def collect():
            results = []
            async for sse in emitter.stream():
                results.append(sse)
            return results

        results = asyncio.run(collect())
        assert len(results) == 1
        assert "file_completed" in results[0]

    def test_multiple_events(self):
        emitter = QueueEmitter()
        emitter.emit(FileStarted(file_path="a", file_index=0, total_files=2, modality="text", filetype="txt"))
        emitter.emit(FileCompleted(file_path="a", records_emitted=1, duration_sec=0.1))
        emitter.close()

        async def collect():
            results = []
            async for sse in emitter.stream():
                results.append(sse)
            return results

        results = asyncio.run(collect())
        assert len(results) == 2

    def test_close_without_events(self):
        emitter = QueueEmitter()
        emitter.close()

        async def collect():
            results = []
            async for sse in emitter.stream():
                results.append(sse)
            return results

        results = asyncio.run(collect())
        assert len(results) == 0