#!/usr/bin/env python3
"""
Optimized batch processing of PDF pages using docling VLM pipeline.

This script loads the model ONCE and processes multiple pages, avoiding
the ~150 second model initialization overhead per page.

Usage:
    python scripts/batch_process_pages.py <input_dir> [output_dir]

Example:
    python scripts/batch_process_pages.py documents/Q1FY25_pages/ outputs/
"""

import sys
import time
from pathlib import Path

from docling.datamodel.base_models import InputFormat
from docling.datamodel.pipeline_options import VlmPipelineOptions
from docling.document_converter import DocumentConverter, PdfFormatOption
from docling.pipeline.vlm_pipeline import VlmPipeline
from loguru import logger


def batch_process_pages(input_dir: Path, output_dir: Path, batch_size: int = 8) -> None:
    """
    Process multiple PDF pages efficiently by reusing the VLM model.

    Args:
        input_dir: Directory containing individual PDF page files
        output_dir: Directory to write output markdown files
        batch_size: Number of pages to process concurrently
    """
    # Create output directory
    output_dir.mkdir(parents=True, exist_ok=True)

    # Find all PDF files
    pdf_files = sorted(input_dir.glob("*.pdf"))

    if not pdf_files:
        logger.error(f"No PDF files found in {input_dir}")
        return

    logger.info(f"Found {len(pdf_files)} PDF files to process")
    logger.info(f"Batch size: {batch_size}")

    # Configure VLM pipeline options
    pipeline_options = VlmPipelineOptions()
    pipeline_options.do_table_structure = True
    pipeline_options.do_ocr = True

    # Initialize converter with VLM pipeline (loads model ONCE)
    logger.info("Initializing DocumentConverter with VLM pipeline...")
    start_init = time.time()

    converter = DocumentConverter(
        format_options={
            InputFormat.PDF: PdfFormatOption(
                pipeline_cls=VlmPipeline,
                pipeline_options=pipeline_options,
            )
        }
    )

    init_time = time.time() - start_init
    logger.info(f"Model initialized in {init_time:.2f} seconds")

    # Process in batches
    total_start = time.time()
    processed = 0

    for i in range(0, len(pdf_files), batch_size):
        batch = pdf_files[i : i + batch_size]
        batch_num = i // batch_size + 1
        total_batches = (len(pdf_files) + batch_size - 1) // batch_size

        logger.info(f"\nProcessing batch {batch_num}/{total_batches} ({len(batch)} files)")
        batch_start = time.time()

        # Convert batch (model is reused across all pages)
        results = converter.convert_all(batch)

        # Write outputs
        for pdf_file, result in zip(batch, results, strict=False):
            output_file = output_dir / f"{pdf_file.stem}.md"
            output_file.write_text(result.document.export_to_markdown())
            logger.info(f"  ✓ {pdf_file.name} -> {output_file.name}")
            processed += 1

        batch_time = time.time() - batch_start
        pages_per_sec = len(batch) / batch_time
        logger.info(f"Batch completed in {batch_time:.2f}s ({pages_per_sec:.2f} pages/sec)")

    total_time = time.time() - total_start
    avg_time = total_time / processed if processed > 0 else 0

    logger.success(f"\n✓ Processed {processed} pages in {total_time:.2f}s")
    logger.success(f"  Average: {avg_time:.2f}s per page (excluding model init)")
    logger.success(f"  Total with init: {total_time + init_time:.2f}s")


def main():
    if len(sys.argv) < 2:
        print(__doc__)
        sys.exit(1)

    input_dir = Path(sys.argv[1])

    if not input_dir.exists():
        logger.error(f"Input directory not found: {input_dir}")
        sys.exit(1)

    # Determine output directory
    if len(sys.argv) >= 3:
        output_dir = Path(sys.argv[2])
    else:
        output_dir = Path("outputs") / input_dir.name

    # Process with optimal batch size for RTX 3090 (24GB VRAM)
    # Adjust batch_size based on VRAM usage
    batch_size = 8  # Start conservative, can increase if VRAM allows

    batch_process_pages(input_dir, output_dir, batch_size=batch_size)


if __name__ == "__main__":
    main()
