Cognitive-rag / 01_data_ingestion / ingestion / events.py
events.py
Raw
"""SSE event types and emitter implementations for ingestion progress streaming."""

from __future__ import annotations

import asyncio
import json
import queue
import time
from dataclasses import asdict, dataclass, field
from typing import Any, Dict, List, Optional, Protocol, runtime_checkable


# ---------------------------------------------------------------------------
# Event dataclasses
# ---------------------------------------------------------------------------

@dataclass(frozen=True)
class IngestionStarted:
    event_type: str = field(default="ingestion_started", init=False)
    job_id: str = ""
    total_files: int = 0
    skipped_files: List[Dict[str, str]] = field(default_factory=list)


@dataclass(frozen=True)
class FileStarted:
    event_type: str = field(default="file_started", init=False)
    file_path: str = ""
    file_index: int = 0
    total_files: int = 0
    modality: str = ""
    filetype: str = ""


@dataclass(frozen=True)
class FileProgress:
    event_type: str = field(default="file_progress", init=False)
    file_path: str = ""
    detail: str = ""
    percent: Optional[float] = None


@dataclass(frozen=True)
class FileCompleted:
    event_type: str = field(default="file_completed", init=False)
    file_path: str = ""
    records_emitted: int = 0
    duration_sec: float = 0.0


@dataclass(frozen=True)
class FileFailed:
    event_type: str = field(default="file_failed", init=False)
    file_path: str = ""
    error: str = ""
    modality: str = ""


@dataclass(frozen=True)
class FileSkipped:
    event_type: str = field(default="file_skipped", init=False)
    file_path: str = ""
    reason: str = ""


@dataclass(frozen=True)
class FileEmpty:
    event_type: str = field(default="file_empty", init=False)
    file_path: str = ""
    modality: str = ""
    reason: str = ""


@dataclass(frozen=True)
class IngestionProgress:
    event_type: str = field(default="ingestion_progress", init=False)
    files_completed: int = 0
    total_files: int = 0
    records_emitted: int = 0
    elapsed_sec: float = 0.0
    eta_sec: Optional[float] = None


@dataclass(frozen=True)
class IngestionCompleted:
    event_type: str = field(default="ingestion_completed", init=False)
    stats: Dict[str, Any] = field(default_factory=dict)


# Union type for all events
SSEEvent = (
    IngestionStarted
    | FileStarted
    | FileProgress
    | FileCompleted
    | FileFailed
    | FileSkipped
    | FileEmpty
    | IngestionProgress
    | IngestionCompleted
)

_SENTINEL = object()


def event_to_dict(event: SSEEvent) -> Dict[str, Any]:
    """Convert an event dataclass to a JSON-serializable dict."""
    return asdict(event)


def event_to_sse(event: SSEEvent) -> str:
    """Format an event as an SSE text block: ``event: <type>\\ndata: <json>\\n\\n``."""
    data = json.dumps(event_to_dict(event), ensure_ascii=True)
    return f"event: {event.event_type}\ndata: {data}\n\n"


# ---------------------------------------------------------------------------
# Emitter protocol + implementations
# ---------------------------------------------------------------------------

@runtime_checkable
class EventEmitter(Protocol):
    def emit(self, event: SSEEvent) -> None: ...


class NullEmitter:
    """No-op emitter for CLI backward-compatibility or tests."""

    def emit(self, event: SSEEvent) -> None:
        pass


class LogEmitter:
    """Prints progress to stdout for CLI usage."""

    def emit(self, event: SSEEvent) -> None:
        etype = event.event_type
        if etype == "ingestion_started":
            assert isinstance(event, IngestionStarted)
            print(f"[ingest] Starting ingestion: {event.total_files} files")
            if event.skipped_files:
                print(f"[ingest] Skipped {len(event.skipped_files)} unsupported files")
        elif etype == "file_started":
            assert isinstance(event, FileStarted)
            print(
                f"[ingest] ({event.file_index + 1}/{event.total_files}) "
                f"Processing {event.file_path} [{event.modality}]"
            )
        elif etype == "file_progress":
            assert isinstance(event, FileProgress)
            pct = f" ({event.percent:.0f}%)" if event.percent is not None else ""
            print(f"[ingest]   {event.detail}{pct}")
        elif etype == "file_completed":
            assert isinstance(event, FileCompleted)
            print(
                f"[ingest]   Done: {event.records_emitted} records "
                f"in {event.duration_sec:.1f}s"
            )
        elif etype == "file_failed":
            assert isinstance(event, FileFailed)
            print(f"[ingest]   FAILED: {event.error}")
        elif etype == "file_skipped":
            assert isinstance(event, FileSkipped)
            print(f"[ingest] Skipped {event.file_path}: {event.reason}")
        elif etype == "file_empty":
            assert isinstance(event, FileEmpty)
            print(f"[ingest]   Empty: {event.reason}")
        elif etype == "ingestion_progress":
            assert isinstance(event, IngestionProgress)
            eta = f", ETA {event.eta_sec:.0f}s" if event.eta_sec is not None else ""
            print(
                f"[ingest] Progress: {event.files_completed}/{event.total_files} files, "
                f"{event.records_emitted} records, {event.elapsed_sec:.1f}s elapsed{eta}"
            )
        elif etype == "ingestion_completed":
            assert isinstance(event, IngestionCompleted)
            counts = event.stats.get("counts", {})
            print(
                f"[ingest] Completed: {counts.get('records_emitted', 0)} records "
                f"from {counts.get('files_succeeded', 0)} files"
            )


class QueueEmitter:
    """Thread-safe emitter that bridges worker threads to an async SSE stream.

    Worker threads call ``emit()`` which puts events on a ``queue.Queue``.
    The async ``stream()`` method yields those events for ``EventSourceResponse``.
    """

    def __init__(self) -> None:
        self._queue: queue.Queue[SSEEvent | object] = queue.Queue()
        self._done = False

    def emit(self, event: SSEEvent) -> None:
        self._queue.put(event)

    def close(self) -> None:
        """Signal that no more events will be emitted."""
        self._done = True
        self._queue.put(_SENTINEL)

    async def stream(self):
        """Async generator that yields SSE-formatted strings.

        Polls the queue in a non-blocking way so the async event loop
        stays responsive. Stops when ``close()`` is called and the
        queue is drained.
        """
        while True:
            try:
                event = self._queue.get_nowait()
            except queue.Empty:
                if self._done:
                    return
                await asyncio.sleep(0.05)
                continue

            if event is _SENTINEL:
                return

            yield event_to_sse(event)


class CallbackEmitter:
    """Emitter that calls a user-supplied function for each event."""

    def __init__(self, callback) -> None:
        self._callback = callback

    def emit(self, event: SSEEvent) -> None:
        self._callback(event)


class MultiEmitter:
    """Broadcasts events to multiple emitters."""

    def __init__(self, *emitters: EventEmitter) -> None:
        self._emitters = emitters

    def emit(self, event: SSEEvent) -> None:
        for emitter in self._emitters:
            emitter.emit(event)