Cognitive-rag / 01_data_ingestion / test_ingestion.py
test_ingestion.py
Raw
#!/usr/bin/env python
"""Automated test suite for the multimodal ingestion pipeline.

Runs per-modality tests against 00_Data_test/ datasets, validates output
schema, and reports clear pass/fail results.

Usage:
    python test_ingestion.py --data-dir 00_Data_test --modality all
    python test_ingestion.py --modality text --verbose
    python test_ingestion.py --modality image
"""

from __future__ import annotations

import argparse
import json
import shutil
import sys
import tempfile
import time
from pathlib import Path
from typing import Dict, List, Optional, Tuple

# Ensure sibling module is importable
sys.path.insert(0, str(Path(__file__).resolve().parent))

from ingest_multimodal import ingest, detect_file_type, iter_files

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

REQUIRED_FIELDS = {"doc_id", "page", "section", "section_index", "content", "source"}
REQUIRED_SOURCE_FIELDS = {"path", "filename", "filetype", "modality"}


def read_jsonl(path: Path) -> List[Dict]:
    records = []
    if not path.exists():
        return records
    with open(path, "r", encoding="utf-8") as f:
        for line in f:
            line = line.strip()
            if line:
                records.append(json.loads(line))
    return records


def validate_record(record: Dict) -> List[str]:
    """Return a list of schema violation messages for a single record."""
    issues = []
    missing = REQUIRED_FIELDS - set(record)
    if missing:
        issues.append(f"Missing fields: {sorted(missing)}")
    if not record.get("doc_id"):
        issues.append("Empty doc_id")
    if not record.get("content"):
        issues.append("Empty content")
    if not record.get("section"):
        issues.append("Empty section")
    source = record.get("source")
    if not isinstance(source, dict):
        issues.append("source is not a dict")
    else:
        for field in REQUIRED_SOURCE_FIELDS:
            if not source.get(field):
                issues.append(f"Missing source.{field}")
        # Video-specific validation
        modality = source.get("modality")
        if modality in ("video", "video_frame"):
            if source.get("timestamp_start") is None:
                issues.append("video record missing timestamp_start")
        if modality == "video_frame":
            if not source.get("frame_path"):
                issues.append("frame record missing frame_path")
    return issues


def has_faster_whisper() -> bool:
    try:
        import faster_whisper  # noqa: F401
        return True
    except Exception:
        return False


def has_ffmpeg() -> bool:
    return shutil.which("ffmpeg") is not None


def has_tesseract() -> bool:
    return shutil.which("tesseract") is not None


def run_ingest(
    input_dir: Path,
    max_files: Optional[int] = None,
    whisper_model: str = "tiny",
    chunk_duration: float = 600.0,
    extract_keyframes: bool = False,
    scene_threshold: float = 0.3,
) -> Tuple[Dict, List[Dict], List[Dict], Path]:
    """Run ingestion into a temp directory and return (stats, records, errors, tmp_dir).

    Note: returns the tmp_dir path so callers can inspect output files.
    The caller is responsible for cleanup if needed (but tempfile handles it).
    """
    tmp_dir_obj = tempfile.mkdtemp()
    tmp = Path(tmp_dir_obj)
    stats = ingest(
        input_dir=input_dir,
        output_dir=tmp,
        output_name="cleaned_documents.jsonl",
        errors_name="ingestion_errors.jsonl",
        pdf_strategy="auto",
        ocr_engine="tesseract",
        ocr_language="eng",
        whisper_model=whisper_model,
        whisper_device="cpu",
        whisper_compute_type="int8",
        max_files=max_files,
        chunk_duration=chunk_duration,
        extract_keyframes=extract_keyframes,
        scene_threshold=scene_threshold,
    )
    records = read_jsonl(tmp / "cleaned_documents.jsonl")
    errors = read_jsonl(tmp / "ingestion_errors.jsonl")
    return stats, records, errors, tmp


# ---------------------------------------------------------------------------
# Individual tests
# ---------------------------------------------------------------------------

def test_text(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test text ingestion (WikiText-103 .txt files)."""
    input_dir = data_dir / "text" / "wikitext-103"
    if not input_dir.exists():
        return "SKIP", "Test data not found. Run download_test_datasets.py --text", None

    stats, records, errors, _tmp = run_ingest(input_dir)
    files = stats["counts"]["files_processed"]
    recs = stats["counts"]["records_emitted"]
    errs = stats["counts"]["errors"]

    if errs > 0:
        return "FAIL", f"{files} files -> {recs} records, {errs} errors", records
    if recs == 0:
        return "FAIL", f"{files} files -> 0 records (expected > 0)", records

    # Verify modality
    for rec in records:
        src = rec.get("source", {})
        if src.get("modality") != "text":
            return "FAIL", f"Expected modality 'text', got '{src.get('modality')}'", records

    return "PASS", f"{files} files -> {recs} records, {errs} errors", records


def test_image(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test image OCR ingestion (CORD-v2 receipt images)."""
    input_dir = data_dir / "image" / "cord-v2" / "images"
    if not input_dir.exists():
        return "SKIP", "Test data not found. Run download_test_datasets.py --image", None

    if not has_tesseract():
        return "SKIP", "tesseract not installed", None

    stats, records, errors, _tmp = run_ingest(input_dir, max_files=5)
    files = stats["counts"]["files_processed"]
    recs = stats["counts"]["records_emitted"]
    n_failed = stats["counts"]["files_failed"]
    n_empty = stats["counts"]["files_empty"]

    if n_failed > 0:
        return "FAIL", f"{files} files -> {recs} records, {n_failed} failures", records
    # Some images may produce empty OCR — that's expected and tracked in empty_files
    if recs == 0:
        return "FAIL", f"{files} files -> 0 records (expected >= 1)", records

    for rec in records:
        src = rec.get("source", {})
        if src.get("modality") != "image":
            return "FAIL", f"Expected modality 'image', got '{src.get('modality')}'", records
        if not src.get("ocr_engine"):
            return "FAIL", "Missing source.ocr_engine on image record", records

    detail = f"{files} files -> {recs} records, {n_empty} empty OCR"
    return "PASS", detail, records


def test_audio(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test audio transcription (LibriSpeech FLAC files)."""
    input_dir = data_dir / "audio_subset"
    if not input_dir.exists():
        return "SKIP", "Test data not found. Run download_test_datasets.py --audio", None

    if not has_faster_whisper():
        return "SKIP", "faster-whisper not available (requires Python <3.13)", None

    stats, records, errors, _tmp = run_ingest(input_dir)
    files = stats["counts"]["files_processed"]
    recs = stats["counts"]["records_emitted"]
    errs = stats["counts"]["errors"]

    if errs > 0:
        return "FAIL", f"{files} files -> {recs} records, {errs} errors", records
    if recs == 0:
        return "FAIL", f"{files} files -> 0 records (expected > 0)", records

    for rec in records:
        src = rec.get("source", {})
        if src.get("modality") != "audio":
            return "FAIL", f"Expected modality 'audio', got '{src.get('modality')}'", records
        if src.get("timestamp_start") is None:
            return "FAIL", "Missing source.timestamp_start on audio record", records

    return "PASS", f"{files} files -> {recs} records, {errs} errors", records


def test_video(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test video transcription (YouTube .webm files)."""
    input_dir = data_dir / "video" / "yt-ugc-subset"
    if not input_dir.exists():
        return "SKIP", "Test data not found. Run download_test_datasets.py --video", None

    if not has_faster_whisper():
        return "SKIP", "faster-whisper not available (requires Python <3.13)", None
    if not has_ffmpeg():
        return "SKIP", "ffmpeg not installed", None

    stats, records, errors, _tmp = run_ingest(input_dir, max_files=2)
    files = stats["counts"]["files_processed"]
    recs = stats["counts"]["records_emitted"]
    errs = stats["counts"]["errors"]

    if errs > 0:
        return "FAIL", f"{files} files -> {recs} records, {errs} errors", records
    if recs == 0:
        return "FAIL", f"{files} files -> 0 records (expected > 0)", records

    for rec in records:
        src = rec.get("source", {})
        if src.get("modality") != "video":
            return "FAIL", f"Expected modality 'video', got '{src.get('modality')}'", records
        if src.get("filetype") != "video":
            return "FAIL", f"Expected filetype 'video', got '{src.get('filetype')}'", records

    return "PASS", f"{files} files -> {recs} records, {errs} errors", records


def test_csv(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test CSV ingestion (mtsamples.csv)."""
    # CSV lives in 00_Data, one level up from test data
    csv_dir = data_dir.parent / "00_Data"
    if not csv_dir.exists() or not (csv_dir / "mtsamples.csv").exists():
        return "SKIP", "00_Data/mtsamples.csv not found", None

    stats, records, errors, _tmp = run_ingest(csv_dir, max_files=1)
    files = stats["counts"]["files_processed"]
    recs = stats["counts"]["records_emitted"]
    errs = stats["counts"]["errors"]

    if errs > 0:
        return "FAIL", f"{files} files -> {recs} records, {errs} errors", records
    if recs == 0:
        return "FAIL", f"{files} files -> 0 records (expected > 0)", records

    for rec in records[:5]:  # spot-check first 5
        src = rec.get("source", {})
        if src.get("modality") != "table":
            return "FAIL", f"Expected modality 'table', got '{src.get('modality')}'", records

    return "PASS", f"{files} files -> {recs} records, {errs} errors", records


def test_mixed(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test mixed-modality batch processing across 00_Data_test/."""
    if not data_dir.exists():
        return "SKIP", "Test data directory not found", None

    stats, records, errors, _tmp = run_ingest(data_dir, max_files=15)
    files = stats["counts"]["files_processed"]
    recs = stats["counts"]["records_emitted"]
    errs = stats["counts"]["errors"]

    # In mixed mode, some errors are acceptable (e.g. missing faster-whisper)
    # but we should get at least some records
    if recs == 0:
        return "FAIL", f"{files} files -> 0 records, {errs} errors", records

    # Check that multiple modalities appear
    modalities = {rec.get("source", {}).get("modality") for rec in records}
    modalities.discard(None)

    detail = f"{files} files -> {recs} records, {errs} errors, modalities: {sorted(modalities)}"
    if len(modalities) < 1:
        return "FAIL", f"No modalities detected. {detail}", records

    return "PASS", detail, records


def test_max_files(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test that max_files cap works correctly at different sizes."""
    input_dir = data_dir / "text" / "wikitext-103"
    if not input_dir.exists():
        return "SKIP", "Test data not found", None

    all_records = []
    for cap in [1, 2, 3]:
        stats, records, _, _tmp = run_ingest(input_dir, max_files=cap)
        files = stats["counts"]["files_processed"]
        if files > cap:
            return "FAIL", f"max_files={cap} but processed {files} files", records
        all_records.extend(records)

    # Also test None (no cap) — should process all files
    stats_all, records_all, _, _tmp = run_ingest(input_dir, max_files=None)
    total = stats_all["counts"]["files_processed"]

    return "PASS", f"max_files caps 1,2,3 respected; uncapped processed {total} files", all_records


def test_reporting(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test that failure reporting includes success_files, failed_files, empty_files."""
    input_dir = data_dir / "image" / "cord-v2" / "images"
    if not input_dir.exists():
        return "SKIP", "Test data not found", None
    if not has_tesseract():
        return "SKIP", "tesseract not installed", None

    stats, records, errors, _tmp = run_ingest(input_dir, max_files=5)

    # Verify new stats fields exist
    counts = stats.get("counts", {})
    for field in ("files_succeeded", "files_failed", "files_empty"):
        if field not in counts:
            return "FAIL", f"Missing counts.{field} in stats", records

    # Verify file lists exist
    for field in ("success_files", "failed_files", "empty_files"):
        if field not in stats:
            return "FAIL", f"Missing {field} list in stats", records

    # Verify counts are consistent
    total = counts["files_succeeded"] + counts["files_failed"] + counts["files_empty"]
    if total != counts["files_processed"]:
        return "FAIL", (f"Counts inconsistent: succeeded({counts['files_succeeded']}) + "
                        f"failed({counts['files_failed']}) + empty({counts['files_empty']}) "
                        f"!= processed({counts['files_processed']})"), records

    # Verify empty_files entries have modality and reason
    for entry in stats["empty_files"]:
        if "modality" not in entry or "reason" not in entry:
            return "FAIL", f"empty_files entry missing modality/reason: {entry}", records

    # Verify failed_files entries have modality and reason (if any)
    for entry in stats["failed_files"]:
        if "modality" not in entry or "reason" not in entry:
            return "FAIL", f"failed_files entry missing modality/reason: {entry}", records

    # Verify error log entries include modality
    for err in errors:
        if "modality" not in err:
            return "FAIL", f"Error log entry missing modality: {err}", records

    detail = (f"succeeded={counts['files_succeeded']}, "
              f"failed={counts['files_failed']}, "
              f"empty={counts['files_empty']}")
    return "PASS", detail, records


def test_video_chunked(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test that chunked video processing produces correct timestamps."""
    input_dir = data_dir / "video" / "yt-ugc-subset"
    if not input_dir.exists():
        return "SKIP", "Test data not found. Run download_test_datasets.py --video", None
    if not has_faster_whisper():
        return "SKIP", "faster-whisper not available (requires Python <3.13)", None
    if not has_ffmpeg():
        return "SKIP", "ffmpeg not installed", None

    # Use a very short chunk duration to force multiple chunks
    stats, records, errors, tmp = run_ingest(
        input_dir, max_files=2, chunk_duration=30.0,
    )
    files = stats["counts"]["files_processed"]
    recs = stats["counts"]["records_emitted"]

    if recs == 0:
        return "FAIL", f"{files} files -> 0 records (expected > 0)", records

    # Verify timestamps are monotonically non-decreasing per doc_id
    by_doc: Dict[str, List[float]] = {}
    for rec in records:
        src = rec.get("source", {})
        if src.get("modality") != "video":
            continue
        doc = rec["doc_id"]
        by_doc.setdefault(doc, []).append(src.get("timestamp_start", 0.0))

    for doc, timestamps in by_doc.items():
        for i in range(1, len(timestamps)):
            if timestamps[i] < timestamps[i - 1]:
                return "FAIL", f"Timestamps not monotonic for {doc} at index {i}", records

    # Verify segment numbering is continuous per doc_id
    seg_by_doc: Dict[str, List[str]] = {}
    for rec in records:
        if rec.get("source", {}).get("modality") == "video":
            seg_by_doc.setdefault(rec["doc_id"], []).append(rec["section"])
    for doc, sections in seg_by_doc.items():
        expected = [f"segment_{i:04d}" for i in range(len(sections))]
        if sections != expected:
            return "FAIL", f"Segment numbering not continuous for {doc}", records

    # Verify no checkpoint files remain after success
    checkpoints = list(tmp.glob(".checkpoint_*.json"))
    if checkpoints:
        return "FAIL", f"Checkpoint files not cleaned up: {checkpoints}", records

    # Clean up
    shutil.rmtree(tmp, ignore_errors=True)
    return "PASS", f"{recs} records with monotonic timestamps across chunks", records


def test_video_resume(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test that checkpoint resume works correctly."""
    input_dir = data_dir / "video" / "yt-ugc-subset"
    if not input_dir.exists():
        return "SKIP", "Test data not found. Run download_test_datasets.py --video", None
    if not has_faster_whisper():
        return "SKIP", "faster-whisper not available (requires Python <3.13)", None
    if not has_ffmpeg():
        return "SKIP", "ffmpeg not installed", None

    # Run a full fresh ingestion with short chunks
    _stats1, records1, _, tmp1 = run_ingest(
        input_dir, max_files=1, chunk_duration=30.0,
    )
    shutil.rmtree(tmp1, ignore_errors=True)

    if not records1:
        return "SKIP", "No records produced (video may be too short)", None

    # Run again — should produce same number of records
    _stats2, records2, _, tmp2 = run_ingest(
        input_dir, max_files=1, chunk_duration=30.0,
    )
    shutil.rmtree(tmp2, ignore_errors=True)

    if len(records1) != len(records2):
        return "FAIL", (f"Inconsistent record counts: "
                        f"run1={len(records1)}, run2={len(records2)}"), records2

    return "PASS", f"Consistent output across runs ({len(records1)} records)", records1


def test_keyframes(data_dir: Path, verbose: bool) -> Tuple[str, str, Optional[List[Dict]]]:
    """Test keyframe extraction with OCR produces frame records."""
    input_dir = data_dir / "video" / "yt-ugc-subset"
    if not input_dir.exists():
        return "SKIP", "Test data not found. Run download_test_datasets.py --video", None
    if not has_faster_whisper():
        return "SKIP", "faster-whisper not available (requires Python <3.13)", None
    if not has_ffmpeg():
        return "SKIP", "ffmpeg not installed", None
    if not has_tesseract():
        return "SKIP", "tesseract not installed", None

    _stats, records, _errors, tmp = run_ingest(
        input_dir, max_files=1, extract_keyframes=True, scene_threshold=0.3,
    )

    segment_records = [r for r in records if r["section"].startswith("segment_")]
    frame_records = [r for r in records if r["section"].startswith("frame_")]

    if not segment_records:
        shutil.rmtree(tmp, ignore_errors=True)
        return "FAIL", "No transcript segments produced", records

    # Verify frame record schema if any frames were extracted
    for rec in frame_records:
        src = rec.get("source", {})
        if src.get("modality") != "video_frame":
            shutil.rmtree(tmp, ignore_errors=True)
            return "FAIL", f"Frame record has wrong modality: {src.get('modality')}", records
        if not src.get("frame_path"):
            shutil.rmtree(tmp, ignore_errors=True)
            return "FAIL", "Frame record missing frame_path", records
        if src.get("timestamp_start") is None:
            shutil.rmtree(tmp, ignore_errors=True)
            return "FAIL", "Frame record missing timestamp", records
        if not src.get("ocr_engine"):
            shutil.rmtree(tmp, ignore_errors=True)
            return "FAIL", "Frame record missing ocr_engine", records

    # Check that frames directory was created if frames were extracted
    if frame_records:
        frames_dirs = list(tmp.glob("frames/*"))
        if not frames_dirs:
            shutil.rmtree(tmp, ignore_errors=True)
            return "FAIL", "Frame records exist but no frames/ directory found", records

    shutil.rmtree(tmp, ignore_errors=True)
    detail = f"{len(segment_records)} segments, {len(frame_records)} frames"
    return "PASS", detail, records


def test_schema(all_records: List[Dict], verbose: bool) -> Tuple[str, str, None]:
    """Validate schema compliance across all collected records."""
    if not all_records:
        return "SKIP", "No records to validate", None

    total_issues = 0
    issue_details = []
    for i, rec in enumerate(all_records):
        issues = validate_record(rec)
        if issues:
            total_issues += len(issues)
            if verbose and len(issue_details) < 10:
                issue_details.append(f"  Record {i}: {'; '.join(issues)}")

    if total_issues > 0:
        detail = f"{total_issues} issues across {len(all_records)} records"
        if issue_details:
            detail += "\n" + "\n".join(issue_details)
        return "FAIL", detail, None

    return "PASS", f"All {len(all_records)} records pass schema validation", None


# ---------------------------------------------------------------------------
# Runner
# ---------------------------------------------------------------------------

TESTS = {
    "text": test_text,
    "image": test_image,
    "audio": test_audio,
    "video": test_video,
    "csv": test_csv,
    "mixed": test_mixed,
    "max_files": test_max_files,
    "reporting": test_reporting,
    "video_chunked": test_video_chunked,
    "video_resume": test_video_resume,
    "keyframes": test_keyframes,
}


def run_tests(data_dir: Path, modality: str, verbose: bool) -> int:
    """Run selected tests and return exit code (0=pass, 1=failures)."""
    print("=== Multimodal Ingestion Test Suite ===\n")

    if modality == "all":
        test_names = list(TESTS.keys())
    else:
        test_names = [modality]

    passed = 0
    failed = 0
    skipped = 0
    all_records: List[Dict] = []

    for name in test_names:
        test_fn = TESTS.get(name)
        if test_fn is None:
            print(f"[????] Unknown test: {name}")
            continue

        start = time.time()
        try:
            status, message, records = test_fn(data_dir, verbose)
        except Exception as exc:
            status = "FAIL"
            message = f"Unhandled exception: {exc}"
            records = None
        elapsed = time.time() - start

        if records:
            all_records.extend(records)

        tag = {"PASS": "PASS", "FAIL": "FAIL", "SKIP": "SKIP"}.get(status, "????")
        time_str = f"({elapsed:.1f}s)" if status != "SKIP" else ""
        print(f"[{tag}] {name:<12} {message} {time_str}")

        if verbose and status == "FAIL" and records:
            for rec in records[:3]:
                print(f"       Sample record: doc_id={rec.get('doc_id')}, "
                      f"modality={rec.get('source', {}).get('modality')}")

        if status == "PASS":
            passed += 1
        elif status == "FAIL":
            failed += 1
        else:
            skipped += 1

    # Schema test runs on all collected records
    if modality == "all" and all_records:
        start = time.time()
        status, message, _ = test_schema(all_records, verbose)
        elapsed = time.time() - start
        tag = {"PASS": "PASS", "FAIL": "FAIL", "SKIP": "SKIP"}.get(status, "????")
        print(f"[{tag}] {'schema':<12} {message} ({elapsed:.1f}s)")
        if status == "PASS":
            passed += 1
        elif status == "FAIL":
            failed += 1
        else:
            skipped += 1

    print(f"\nResults: {passed} passed, {failed} failed, {skipped} skipped")
    return 1 if failed > 0 else 0


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Test the multimodal ingestion pipeline.")
    parser.add_argument(
        "--data-dir",
        default="00_Data_test",
        help="Root directory containing test data (default: 00_Data_test).",
    )
    parser.add_argument(
        "--modality",
        choices=["text", "image", "audio", "video", "csv", "mixed", "max_files",
                 "reporting", "video_chunked", "video_resume", "keyframes", "all"],
        default="all",
        help="Which test(s) to run (default: all).",
    )
    parser.add_argument(
        "--verbose",
        action="store_true",
        help="Print detailed output for each test.",
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    data_dir = Path(args.data_dir)
    if not data_dir.exists():
        print(f"Error: data directory '{data_dir}' not found.")
        print("Run download_test_datasets.py first to download test data.")
        sys.exit(1)
    exit_code = run_tests(data_dir, args.modality, args.verbose)
    sys.exit(exit_code)


if __name__ == "__main__":
    main()