#!/usr/bin/env python3
"""
Phase 0: Preprocess scanned PDF pages for better OCR.
Applies bilateral filter → adaptive threshold → morphology → output clean PDF.

Usage:
    python3 preprocess_all.py <input.pdf> [--output clean.pdf] [--dpi 300] [--debug]
    
Config loaded from preprocess-config.yaml (next to this script) or defaults.
"""
import sys
import os
import subprocess
import tempfile
import shutil
import argparse
from pathlib import Path

try:
    import cv2
    import numpy as np
except ImportError:
    print("❌ Need opencv-python: pip install opencv-python")
    sys.exit(1)

# Default config (matches guide: 04-OCR-Pipeline-Full.md)
DEFAULT_CONFIG = {
    "preprocessing": {
        "bilateral_d": 9,
        "bilateral_sigmaColor": 75,
        "bilateral_sigmaSpace": 75,
        "adaptive_threshold": {
            "method": "ADAPTIVE_THRESH_GAUSSIAN_C",
            "block_size": 21,
            "C": 25
        },
        "morphology": {
            "kernel": [3, 3],
            "mode": "MORPH_CLOSE"
        },
        "output": {
            "dpi": 300,
            "format": "png"
        }
    }
}


def load_config():
    """Try loading from YAML, return defaults if not found."""
    config_path = os.path.join(os.path.dirname(__file__), "preprocess-config.yaml")
    if os.path.exists(config_path):
        try:
            import yaml
            with open(config_path) as f:
                return yaml.safe_load(f)
        except Exception:
            pass
    return DEFAULT_CONFIG


def process_page(input_path, output_path, cfg, debug=False):
    """Apply preprocessing to a single page image."""
    pp = cfg["preprocessing"]
    
    # Read image
    img = cv2.imread(input_path, cv2.IMREAD_GRAYSCALE)
    if img is None:
        print(f"  ⚠️  Cannot read: {input_path}")
        return False
    
    # 1. Bilateral filter (preserves edges while reducing noise)
    d = pp["bilateral_d"]
    sc = pp["bilateral_sigmaColor"]
    ss = pp["bilateral_sigmaSpace"]
    img = cv2.bilateralFilter(img, d, sc, ss)
    
    if debug:
        cv2.imwrite(output_path.replace(".png", "_1_bilateral.png"), img)
    
    # 2. Adaptive threshold (binarization)
    at = pp["adaptive_threshold"]
    method_map = {
        "ADAPTIVE_THRESH_GAUSSIAN_C": cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
        "ADAPTIVE_THRESH_MEAN_C": cv2.ADAPTIVE_THRESH_MEAN_C,
    }
    method = method_map.get(at["method"], cv2.ADAPTIVE_THRESH_GAUSSIAN_C)
    block = at["block_size"]
    C = at["C"]
    # block_size must be odd
    if block % 2 == 0:
        block += 1
    img = cv2.adaptiveThreshold(img, 255, method, cv2.THRESH_BINARY, block, C)
    
    if debug:
        cv2.imwrite(output_path.replace(".png", "_2_threshold.png"), img)
    
    # 3. Morphology (close small holes, remove noise)
    morph = pp["morphology"]
    kernel = np.ones(tuple(morph["kernel"]), np.uint8)
    mode_map = {
        "MORPH_CLOSE": cv2.MORPH_CLOSE,
        "MORPH_OPEN": cv2.MORPH_OPEN,
        "MORPH_DILATE": cv2.MORPH_DILATE,
        "MORPH_ERODE": cv2.MORPH_ERODE,
    }
    mode = mode_map.get(morph["mode"], cv2.MORPH_CLOSE)
    img = cv2.morphologyEx(img, mode, kernel)
    
    if debug:
        cv2.imwrite(output_path.replace(".png", "_3_morph.png"), img)
    
    # Save
    cv2.imwrite(output_path, img)
    return True


def extract_pages(pdf_path, work_dir, dpi=300):
    """Extract PDF pages to PNG using pdftoppm."""
    prefix = os.path.join(work_dir, "page")
    cmd = ["pdftoppm", "-r", str(dpi), "-png", pdf_path, prefix]
    print(f"  📄 Extracting: {' '.join(cmd)}")
    result = subprocess.run(cmd, capture_output=True, text=True)
    if result.returncode != 0:
        print(f"  ❌ pdftoppm failed: {result.stderr}")
        return []
    
    pages = sorted(Path(work_dir).glob("page-*.png"))
    print(f"  ✅ Extracted {len(pages)} pages")
    return pages


def build_clean_pdf(page_dir, output_path):
    """Merge processed PNGs back into a PDF using img2pdf."""
    pages = sorted(Path(page_dir).glob("page-*.png"))
    if not pages:
        print("  ❌ No pages to merge")
        return False
    
    cmd = ["img2pdf"] + [str(p) for p in pages] + ["-o", output_path]
    print(f"  📕 Merging {len(pages)} pages → {output_path}")
    result = subprocess.run(cmd, capture_output=True, text=True)
    if result.returncode != 0:
        print(f"  ❌ img2pdf failed: {result.stderr}")
        return False
    
    size_mb = os.path.getsize(output_path) / (1024 * 1024)
    print(f"  ✅ Clean PDF: {size_mb:.1f} MB")
    return True


def main():
    parser = argparse.ArgumentParser(description="Phase 0: Preprocess scanned PDF for OCR")
    parser.add_argument("input", help="Input PDF path")
    parser.add_argument("--output", "-o", help="Output clean PDF path (default: input-clean.pdf)")
    parser.add_argument("--dpi", type=int, default=300, help="DPI for extraction (default: 300)")
    parser.add_argument("--debug", action="store_true", help="Save intermediate debug images")
    args = parser.parse_args()
    
    if not os.path.exists(args.input):
        print(f"❌ File not found: {args.input}")
        sys.exit(1)
    
    input_path = os.path.abspath(args.input)
    base = os.path.splitext(os.path.basename(input_path))[0]
    output_path = args.output or os.path.join(os.path.dirname(input_path), f"{base}-clean.pdf")
    output_path = os.path.abspath(output_path)
    
    cfg = load_config()
    pp = cfg["preprocessing"]
    
    print("=" * 60)
    print(f"🧹 Phase 0: Preprocessing")
    print(f"   Input:  {input_path}")
    print(f"   Output: {output_path}")
    print(f"   Config: bilateral(d={pp['bilateral_d']}, sc={pp['bilateral_sigmaColor']}, ss={pp['bilateral_sigmaSpace']})")
    print(f"           threshold({pp['adaptive_threshold']['method']}, block={pp['adaptive_threshold']['block_size']}, C={pp['adaptive_threshold']['C']})")
    print(f"           morph({pp['morphology']['mode']}, kernel={pp['morphology']['kernel']})")
    print("=" * 60)
    
    # Create temp workspace
    work_dir = tempfile.mkdtemp(prefix="ocr_preprocess_")
    try:
        # Step 1: Extract pages
        print("\n📤 Step 1/3: Extract PDF → PNG")
        pages = extract_pages(input_path, work_dir, dpi=args.dpi)
        if not pages:
            sys.exit(1)
        
        # Step 2: Process each page
        print(f"\n🔧 Step 2/3: Process {len(pages)} pages")
        cleaned_dir = os.path.join(work_dir, "cleaned")
        os.makedirs(cleaned_dir, exist_ok=True)
        
        for i, page_path in enumerate(pages):
            fname = os.path.basename(page_path)
            out_path = os.path.join(cleaned_dir, fname)
            pct = (i + 1) / len(pages) * 100
            print(f"  [{i+1:03d}/{len(pages)} {pct:5.1f}%] {fname}...", end=" ", flush=True)
            process_page(str(page_path), out_path, cfg, debug=args.debug)
            size_kb = os.path.getsize(out_path) / 1024
            print(f"✅ {size_kb:.0f}KB")
        
        # Step 3: Merge back
        print(f"\n📕 Step 3/3: Merge → Clean PDF")
        build_clean_pdf(cleaned_dir, output_path)
        
        print(f"\n🎉 Done! Clean PDF: {output_path}")
        
    finally:
        shutil.rmtree(work_dir, ignore_errors=True)


if __name__ == "__main__":
    main()
