Cognitive-rag / 01_data_ingestion / ingestion / engine.py
engine.py
Raw
"""Parallel ingestion engine with event-driven progress reporting."""

from __future__ import annotations

import json
import threading
import time
from collections import Counter
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List, Optional

import sys, os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))

from ingest_multimodal import (
    OcrRunner,
    WhisperTranscriber,
    add_relations,
    build_doc_id,
    detect_file_type,
    iter_files,
    load_langchain_documents,
    process_video_chunked,
    records_from_audio,
    records_from_csv,
    records_from_documents,
    records_from_image,
    _resolve_modality,
)
from ingestion.events import (
    EventEmitter,
    FileCompleted,
    FileEmpty,
    FileFailed,
    FileProgress,
    FileSkipped,
    FileStarted,
    IngestionCompleted,
    IngestionProgress,
    IngestionStarted,
    NullEmitter,
)


@dataclass
class IngestionConfig:
    """All configuration for an ingestion run."""

    input_dir: Path
    output_dir: Path
    output_name: str = "cleaned_documents.jsonl"
    errors_name: str = "ingestion_errors.jsonl"
    pdf_strategy: str = "auto"
    ocr_engine: str = "tesseract"
    ocr_language: str = "eng"
    whisper_model: str = "small"
    whisper_device: str = "cpu"
    whisper_compute_type: str = "int8"
    max_files: Optional[int] = None
    chunk_duration: float = 600.0
    extract_keyframes: bool = False
    scene_threshold: float = 0.3
    workers: int = 4


@dataclass
class FileResult:
    """Result from processing a single file."""

    file_index: int
    path: Path
    filetype: str
    modality: str
    records: List[Dict[str, Any]] = field(default_factory=list)
    error: Optional[str] = None
    empty: bool = False
    duration_sec: float = 0.0


class CancelledError(Exception):
    pass


class IngestionEngine:
    """Concurrent ingestion engine using ThreadPoolExecutor.

    - PDF / text / CSV / image work runs freely in parallel threads
    - Audio / video (Whisper) is serialized via ``_whisper_lock`` to avoid
      model contention (safe on both CPU and GPU)
    - Cooperative cancellation via ``threading.Event``
    - Emits SSE events for every lifecycle transition
    """

    def __init__(
        self,
        config: IngestionConfig,
        emitter: EventEmitter | None = None,
    ) -> None:
        self.config = config
        self.emitter: EventEmitter = emitter or NullEmitter()
        self._cancel = threading.Event()
        self._whisper_lock = threading.Lock()
        self._transcriber: Optional[WhisperTranscriber] = None
        self._ocr_runner = OcrRunner(config.ocr_engine, config.ocr_language)

        # Thread-safe progress counters
        self._lock = threading.Lock()
        self._files_completed = 0
        self._total_records = 0
        self._start_time: float = 0.0

    # ------------------------------------------------------------------
    # Public API
    # ------------------------------------------------------------------

    def run(self) -> Dict[str, Any]:
        """Execute the ingestion pipeline. Returns stats dict."""
        self.config.output_dir.mkdir(parents=True, exist_ok=True)
        self._start_time = time.monotonic()

        # Discover and classify files
        all_files = list(iter_files(self.config.input_dir, self.config.max_files))
        supported, skipped = self._classify_files(all_files)

        self.emitter.emit(IngestionStarted(
            total_files=len(supported),
            skipped_files=skipped,
        ))

        # Emit individual skip events
        for s in skipped:
            self.emitter.emit(FileSkipped(file_path=s["path"], reason=s["reason"]))

        # Process files
        if self.config.workers <= 0 or len(supported) == 0:
            results = self._run_sequential(supported)
        else:
            results = self._run_parallel(supported)

        # Sort by original file index for deterministic output
        results.sort(key=lambda r: r.file_index)

        # Write output and build stats
        stats = self._write_output(results, skipped)
        self.emitter.emit(IngestionCompleted(stats=stats))
        return stats

    def cancel(self) -> None:
        """Signal all workers to stop after their current file."""
        self._cancel.set()

    @property
    def is_cancelled(self) -> bool:
        return self._cancel.is_set()

    # ------------------------------------------------------------------
    # File classification
    # ------------------------------------------------------------------

    def _classify_files(
        self, files: List[Path]
    ) -> tuple[List[tuple[int, Path, str, str]], List[Dict[str, str]]]:
        """Split files into supported (with index, path, filetype, modality)
        and skipped (with path + reason)."""
        supported: List[tuple[int, Path, str, str]] = []
        skipped: List[Dict[str, str]] = []
        idx = 0
        for path in files:
            filetype = detect_file_type(path)
            if not filetype:
                rel = self._rel(path)
                skipped.append({
                    "path": rel,
                    "reason": f"Unsupported file type: {path.suffix}",
                })
                continue
            modality = _resolve_modality(filetype)
            supported.append((idx, path, filetype, modality))
            idx += 1
        return supported, skipped

    # ------------------------------------------------------------------
    # Execution strategies
    # ------------------------------------------------------------------

    def _run_sequential(
        self, supported: List[tuple[int, Path, str, str]]
    ) -> List[FileResult]:
        results: List[FileResult] = []
        for idx, path, filetype, modality in supported:
            if self._cancel.is_set():
                break
            result = self._process_file(idx, path, filetype, modality, len(supported))
            results.append(result)
            self._update_progress(result, len(supported))
        return results

    def _run_parallel(
        self, supported: List[tuple[int, Path, str, str]]
    ) -> List[FileResult]:
        results: List[FileResult] = []
        total = len(supported)
        with ThreadPoolExecutor(max_workers=self.config.workers) as pool:
            futures = {
                pool.submit(
                    self._process_file, idx, path, filetype, modality, total
                ): idx
                for idx, path, filetype, modality in supported
            }
            for future in as_completed(futures):
                if self._cancel.is_set():
                    pool.shutdown(wait=False, cancel_futures=True)
                    break
                result = future.result()
                results.append(result)
                self._update_progress(result, total)
        return results

    # ------------------------------------------------------------------
    # Per-file processing
    # ------------------------------------------------------------------

    def _process_file(
        self,
        file_index: int,
        path: Path,
        filetype: str,
        modality: str,
        total_files: int,
    ) -> FileResult:
        if self._cancel.is_set():
            return FileResult(
                file_index=file_index,
                path=path,
                filetype=filetype,
                modality=modality,
                error="Cancelled",
            )

        rel_path = self._rel(path)
        self.emitter.emit(FileStarted(
            file_path=rel_path,
            file_index=file_index,
            total_files=total_files,
            modality=modality,
            filetype=filetype,
        ))

        t0 = time.monotonic()
        doc_id = build_doc_id(path, self.config.input_dir)

        try:
            records = self._extract_records(path, filetype, modality, doc_id, rel_path)
        except Exception as exc:
            duration = time.monotonic() - t0
            self.emitter.emit(FileFailed(
                file_path=rel_path, error=str(exc), modality=modality
            ))
            return FileResult(
                file_index=file_index,
                path=path,
                filetype=filetype,
                modality=modality,
                error=str(exc),
                duration_sec=duration,
            )

        duration = time.monotonic() - t0

        if not records:
            self.emitter.emit(FileEmpty(
                file_path=rel_path,
                modality=modality,
                reason="Extraction produced no content",
            ))
            return FileResult(
                file_index=file_index,
                path=path,
                filetype=filetype,
                modality=modality,
                empty=True,
                duration_sec=duration,
            )

        self.emitter.emit(FileCompleted(
            file_path=rel_path,
            records_emitted=len(records),
            duration_sec=duration,
        ))
        return FileResult(
            file_index=file_index,
            path=path,
            filetype=filetype,
            modality=modality,
            records=records,
            duration_sec=duration,
        )

    def _extract_records(
        self,
        path: Path,
        filetype: str,
        modality: str,
        doc_id: str,
        rel_path: str,
    ) -> List[Dict[str, Any]]:
        """Route to the correct extraction function.

        Audio and video acquire the whisper lock; everything else runs freely.
        """
        root = self.config.input_dir

        if filetype == "csv":
            records = records_from_csv(path, doc_id, root)
            add_relations(records)
            return records

        if filetype == "image":
            records = records_from_image(path, doc_id, root, self._ocr_runner)
            add_relations(records)
            return records

        if filetype == "audio":
            with self._whisper_lock:
                self._ensure_transcriber()
                records = records_from_audio(
                    path, doc_id, root, self._transcriber, "audio"
                )
            add_relations(records)
            return records

        if filetype == "video":
            total_chunks = self._estimate_video_chunks(path)
            with self._whisper_lock:
                self._ensure_transcriber()
                records = self._process_video_with_progress(
                    path, doc_id, rel_path, total_chunks
                )
            return records  # relations already added by process_video_chunked

        # Text types (pdf, docx, pptx, txt, md, html)
        docs, loader_name = load_langchain_documents(
            path, filetype, self.config.pdf_strategy
        )
        records = records_from_documents(
            docs, doc_id, path, root, filetype, "text", loader_name
        )
        add_relations(records)
        return records

    def _process_video_with_progress(
        self,
        path: Path,
        doc_id: str,
        rel_path: str,
        total_chunks: int,
    ) -> List[Dict[str, Any]]:
        """Wrap process_video_chunked with chunk-level progress events."""
        self.emitter.emit(FileProgress(
            file_path=rel_path,
            detail=f"Video: {total_chunks} chunks to process",
            percent=0.0,
        ))
        records = process_video_chunked(
            video_path=path,
            doc_id=doc_id,
            root=self.config.input_dir,
            transcriber=self._transcriber,
            output_dir=self.config.output_dir,
            chunk_duration=self.config.chunk_duration,
            extract_keyframes=self.config.extract_keyframes,
            ocr_runner=self._ocr_runner if self.config.extract_keyframes else None,
            scene_threshold=self.config.scene_threshold,
        )
        return records

    def _estimate_video_chunks(self, path: Path) -> int:
        """Get approximate chunk count without holding the whisper lock."""
        try:
            from ingest_multimodal import get_video_duration
            duration = get_video_duration(path)
            cd = self.config.chunk_duration
            return max(1, int(duration // cd) + (1 if duration % cd > 0 else 0))
        except Exception:
            return 1

    # ------------------------------------------------------------------
    # Progress tracking
    # ------------------------------------------------------------------

    def _update_progress(self, result: FileResult, total_files: int) -> None:
        with self._lock:
            self._files_completed += 1
            self._total_records += len(result.records)
            elapsed = time.monotonic() - self._start_time
            completed = self._files_completed
            records = self._total_records

        # ETA: simple linear extrapolation
        eta = None
        if completed > 0 and completed < total_files:
            avg_per_file = elapsed / completed
            remaining = total_files - completed
            eta = avg_per_file * remaining

        self.emitter.emit(IngestionProgress(
            files_completed=completed,
            total_files=total_files,
            records_emitted=records,
            elapsed_sec=round(elapsed, 2),
            eta_sec=round(eta, 2) if eta is not None else None,
        ))

    # ------------------------------------------------------------------
    # Output writing
    # ------------------------------------------------------------------

    def _write_output(
        self,
        results: List[FileResult],
        skipped: List[Dict[str, str]],
    ) -> Dict[str, Any]:
        output_path = self.config.output_dir / self.config.output_name
        errors_path = self.config.output_dir / self.config.errors_name

        file_counts: Counter = Counter()
        record_counts: Counter = Counter()
        success_files: List[str] = []
        failed_files: List[Dict[str, str]] = []
        empty_files: List[Dict[str, str]] = []
        total_records = 0

        with open(output_path, "w", encoding="utf-8") as out_handle, \
             open(errors_path, "w", encoding="utf-8") as error_handle:
            for result in results:
                file_counts[result.filetype] += 1
                rel = self._rel(result.path)

                if result.error:
                    failed_files.append({
                        "path": rel,
                        "modality": result.modality,
                        "reason": result.error,
                    })
                    error_entry = {
                        "path": rel,
                        "stage": "ingest",
                        "error": result.error,
                        "modality": result.modality,
                        "filetype": result.filetype,
                    }
                    error_handle.write(json.dumps(error_entry, ensure_ascii=True) + "\n")
                elif result.empty:
                    empty_files.append({
                        "path": rel,
                        "modality": result.modality,
                        "reason": "Extraction produced no content",
                    })
                    error_entry = {
                        "path": rel,
                        "stage": "empty_content",
                        "error": "Extraction produced no content",
                        "modality": result.modality,
                        "filetype": result.filetype,
                    }
                    error_handle.write(json.dumps(error_entry, ensure_ascii=True) + "\n")
                else:
                    for record in result.records:
                        out_handle.write(json.dumps(record, ensure_ascii=True) + "\n")
                    total_records += len(result.records)
                    record_counts[result.filetype] += len(result.records)
                    success_files.append(rel)

        elapsed = time.monotonic() - self._start_time
        stats: Dict[str, Any] = {
            "input_root": str(self.config.input_dir),
            "output_file": str(output_path),
            "errors_file": str(errors_path),
            "generated_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"),
            "config": {
                "pdf_strategy": self.config.pdf_strategy,
                "ocr_engine": self.config.ocr_engine,
                "ocr_language": self.config.ocr_language,
                "whisper_model": self.config.whisper_model,
                "whisper_device": self.config.whisper_device,
                "whisper_compute_type": self.config.whisper_compute_type,
                "video_chunk_duration": self.config.chunk_duration,
                "extract_keyframes": self.config.extract_keyframes,
                "scene_threshold": (
                    self.config.scene_threshold if self.config.extract_keyframes else None
                ),
                "workers": self.config.workers,
            },
            "counts": {
                "files_processed": len(results),
                "files_succeeded": len(success_files),
                "files_failed": len(failed_files),
                "files_empty": len(empty_files),
                "files_skipped": len(skipped),
                "records_emitted": total_records,
                "errors": len(failed_files) + len(empty_files),
            },
            "duration_sec": round(elapsed, 2),
            "files_by_type": dict(file_counts),
            "records_by_type": dict(record_counts),
            "success_files": success_files,
            "failed_files": failed_files,
            "empty_files": empty_files,
            "skipped_files": skipped,
        }

        stats_path = self.config.output_dir / "ingestion_stats.json"
        with open(stats_path, "w", encoding="utf-8") as f:
            json.dump(stats, f, indent=2, ensure_ascii=True)

        return stats

    # ------------------------------------------------------------------
    # Helpers
    # ------------------------------------------------------------------

    def _ensure_transcriber(self) -> None:
        """Lazily initialize the Whisper transcriber (must hold _whisper_lock)."""
        if self._transcriber is None:
            self._transcriber = WhisperTranscriber(
                self.config.whisper_model,
                self.config.whisper_device,
                self.config.whisper_compute_type,
            )

    def _rel(self, path: Path) -> str:
        try:
            return path.relative_to(self.config.input_dir).as_posix()
        except ValueError:
            return path.as_posix()