"""
invoice_ocr.py — Facturen uitlezen met Unlimited-OCR en structureren naar JSON.

Gebruik:
    python invoice_ocr.py input/factuur.pdf
    python invoice_ocr.py input/factuur.pdf --output-dir output --model baidu/Unlimited-OCR

Pijplijn:
    PDF -> afbeeldingen -> Unlimited-OCR (markdown) -> LLM-extractie (JSON) -> validatie
"""

from __future__ import annotations

import argparse
import json
import re
import shutil
import sys
import tempfile
from pathlib import Path

# Zware imports (torch/transformers) gebeuren pas in load_model(), zodat
# --help en argumentfouten meteen werken zonder het model te laden.


# --------------------------------------------------------------------------
# 1. PDF -> afbeeldingen
# --------------------------------------------------------------------------
def pdf_to_images(pdf_path: Path, dpi: int = 300) -> tuple[list[str], Path]:
    """Zet elke PDF-pagina om naar een PNG in een tijdelijke map."""
    import fitz  # PyMuPDF

    tmp_dir = Path(tempfile.mkdtemp(prefix="invoice_ocr_"))
    scale = fitz.Matrix(dpi / 72, dpi / 72)

    paths: list[str] = []
    with fitz.open(pdf_path) as doc:
        for i, page in enumerate(doc):
            out = tmp_dir / f"page_{i + 1:04d}.png"
            page.get_pixmap(matrix=scale).save(out)
            paths.append(str(out))
    return paths, tmp_dir


# --------------------------------------------------------------------------
# 2. Unlimited-OCR: afbeeldingen -> markdown
# --------------------------------------------------------------------------
_DET_RE = re.compile(r"<\|det\|>([^<\s]+)(?:\s*\[[^\]]*\])?\s*<\|/det\|>(.*)", re.DOTALL)


def strip_det_markers(raw: str) -> str:
    """Verwijder de <|det|>type [bbox]<|/det|> markers uit de OCR-uitvoer."""
    blocks: list[list[str]] = []
    current: list[str] | None = None
    for line in raw.splitlines():
        line = line.rstrip()
        if not line:
            continue
        m = _DET_RE.match(line)
        if m:
            category, content = m.group(1).strip(), m.group(2).strip()
            if category == "image":
                continue
            if current is not None:
                blocks.append(current)
            current = [content] if content else []
            continue
        if current is None:
            current = []
        current.append(line)
    if current is not None:
        blocks.append(current)
    return "\n\n".join("\n".join(b) for b in blocks).strip()


class OCRModel:
    """Laadt Unlimited-OCR één keer en houdt het in het geheugen."""

    def __init__(self, model_name: str) -> None:
        import torch
        from transformers import AutoModel, AutoTokenizer

        if not torch.cuda.is_available():
            raise RuntimeError(
                "Geen CUDA-GPU gevonden. Draai dit op een NVIDIA-GPU, "
                "of gebruik de online demo (zie de guide)."
            )

        self.tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
        self.model = (
            AutoModel.from_pretrained(
                model_name,
                trust_remote_code=True,
                use_safetensors=True,
                torch_dtype=torch.bfloat16,
            )
            .eval()
            .cuda()
        )

    def read(self, image_paths: list[str]) -> str:
        """Parseert de pagina's en geeft de opgeschoonde markdown terug."""
        out_dir = Path(tempfile.mkdtemp(prefix="invoice_ocr_out_"))
        try:
            self.model.infer_multi(
                self.tokenizer,
                prompt="<image>Multi page parsing.",
                image_files=image_paths,
                output_path=str(out_dir),
                image_size=1024,
                max_length=32768,
                no_repeat_ngram_size=35,
                ngram_window=1024,
                save_results=True,
            )
            # Unlimited-OCR schrijft de resultaten naar output_path. De extensie
            # kan per modelversie verschillen; we lezen alle tekstuele resultaten.
            texts: list[str] = []
            for pattern in ("*.mmd", "*.md", "*.txt"):
                for path in sorted(out_dir.glob(pattern)):
                    texts.append(path.read_text(encoding="utf-8"))
            return strip_det_markers("\n".join(texts))
        finally:
            shutil.rmtree(out_dir, ignore_errors=True)


# --------------------------------------------------------------------------
# 3. Markdown -> gestructureerde JSON (tweede LLM-stap)
# --------------------------------------------------------------------------
EXTRACTION_PROMPT = """Je krijgt de uitgelezen tekst van een factuur.
Geef uitsluitend geldige JSON terug die exact dit schema volgt:

{schema}

Regels:
- Bedragen als getallen (punt als decimaalteken), geen valutasymbool.
- Datums als YYYY-MM-DD.
- Onbekende velden: null. lines is een lijst (mag leeg zijn).
- Geen uitleg, alleen de JSON.

Factuurtekst:
---
{text}
---
"""


def extract_json(markdown: str, model: str = "llama3.1") -> dict:
    """Laat een lokaal LLM (via Ollama) de tekst omzetten naar JSON."""
    import ollama

    from invoice_schema import Invoice

    prompt = EXTRACTION_PROMPT.format(schema=Invoice.model_json_schema(), text=markdown)
    response = ollama.chat(
        model=model,
        messages=[{"role": "user", "content": prompt}],
        format="json",
        options={"temperature": 0},
    )
    return json.loads(response["message"]["content"])


# --------------------------------------------------------------------------
# 4. CLI
# --------------------------------------------------------------------------
def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Lees een PDF-factuur uit met Unlimited-OCR.")
    parser.add_argument("pdf", type=Path, help="Pad naar de PDF-factuur.")
    parser.add_argument("--output-dir", type=Path, default=Path("output"), help="Uitvoermap.")
    parser.add_argument("--model", default="baidu/Unlimited-OCR", help="OCR-model.")
    parser.add_argument("--llm", default="llama3.1", help="Ollama-model voor extractie.")
    parser.add_argument("--dpi", type=int, default=300, help="Resolutie voor PDF-conversie.")
    return parser.parse_args()


def main() -> int:
    from invoice_schema import Invoice, validate_invoice

    args = parse_args()

    if not args.pdf.is_file():
        print(f"Bestand niet gevonden: {args.pdf}", file=sys.stderr)
        return 1

    args.output_dir.mkdir(parents=True, exist_ok=True)
    stem = args.pdf.stem

    images, tmp_dir = pdf_to_images(args.pdf, dpi=args.dpi)
    try:
        print(f"[1/3] {len(images)} pagina('s) uitlezen met {args.model} ...")
        ocr = OCRModel(args.model)
        markdown = ocr.read(images)
        md_path = args.output_dir / f"{stem}.md"
        md_path.write_text(markdown, encoding="utf-8")
        print(f"      Markdown opgeslagen: {md_path}")

        print(f"[2/3] Gegevens structureren met {args.llm} ...")
        data = extract_json(markdown, model=args.llm)
        invoice = Invoice.model_validate(data)

        print("[3/3] Valideren ...")
        problems = validate_invoice(invoice)
        result = invoice.model_dump(mode="json")
        result["_needs_review"] = bool(problems)
        result["_problems"] = problems

        json_path = args.output_dir / f"{stem}.json"
        json_path.write_text(json.dumps(result, indent=2, ensure_ascii=False), encoding="utf-8")
        print(f"      JSON opgeslagen: {json_path}")

        if problems:
            print("\n⚠  Handmatige controle nodig:")
            for p in problems:
                print(f"   - {p}")
        else:
            print("\n✓  Validatie geslaagd. Klaar om te verwerken.")
    finally:
        shutil.rmtree(tmp_dir, ignore_errors=True)

    return 0


if __name__ == "__main__":
    raise SystemExit(main())
