#!/usr/bin/env python3
"""Multi-Modal-RAG PoC — Aria Tier-2 Gap G2-2.

Pattern adoptiert aus BMW Architects-Manual Multi-Modal-RAG-Sektion
(Option 2: Embed Text-Summaries + Link zum Raw-Image/Table).

PoC-Scope:
- Input: PDF mit Tabellen + Bildern
- Extract: PyMuPDF (images) + pdfplumber (tables)
- Summarize: Claude (Bedrock-aequivalent: Anthropic-API direkt)
- Store: Summaries als pgvector/SQLite-FTS, Originals als Files in artifact-bucket
- Retrieve: Summary-Vector-Search → Pull Original-Datei → Pass an LLM

Production-Pfad waere: Supabase pgvector + S3-Object-Store + LangChain MultiVectorRetriever.
Dieser PoC zeigt das Pattern in 200 LOC ohne externe Vector-DB (SQLite-FTS5 Fallback).

Usage:
  python3 aria-multimodal-rag-poc.py ingest <pdf-path>
  python3 aria-multimodal-rag-poc.py query "<question>"

Adopt-Item: F1 (BMW), G2-2 (Aria-Tech-Stack-Gap)
"""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import sqlite3
import sys
from pathlib import Path
from datetime import datetime, timezone

DB_PATH = Path("/root/aria/state/multimodal-rag/index.sqlite")
ARTIFACT_DIR = Path("/root/aria/state/multimodal-rag/artifacts")


def ensure_paths():
    DB_PATH.parent.mkdir(parents=True, exist_ok=True)
    ARTIFACT_DIR.mkdir(parents=True, exist_ok=True)
    con = sqlite3.connect(DB_PATH)
    con.executescript("""
        CREATE TABLE IF NOT EXISTS artifacts (
            id TEXT PRIMARY KEY,
            source_path TEXT,
            artifact_type TEXT,
            page INT,
            artifact_index INT,
            artifact_path TEXT,
            summary TEXT,
            created_at TEXT
        );
        CREATE VIRTUAL TABLE IF NOT EXISTS summaries USING fts5(
            id UNINDEXED, summary, content=artifacts, content_rowid=rowid
        );
        CREATE TRIGGER IF NOT EXISTS artifacts_ai AFTER INSERT ON artifacts BEGIN
            INSERT INTO summaries(rowid, id, summary) VALUES (new.rowid, new.id, new.summary);
        END;
    """)
    con.commit()
    return con


def hash_id(*parts: str) -> str:
    h = hashlib.sha256(" ".join(parts).encode()).hexdigest()
    return h[:16]


def extract_pdf(pdf_path: Path) -> list[dict]:
    """Extract images + tables from PDF using PyMuPDF + pdfplumber.

    Returns list of {type, page, index, content_path, summary_input}.
    """
    artifacts = []
    try:
        import fitz  # PyMuPDF
    except ImportError:
        print("ERROR: PyMuPDF not installed (pip install pymupdf)")
        return []
    try:
        import pdfplumber
    except ImportError:
        print("WARN: pdfplumber not installed — tables skipped")
        pdfplumber = None

    doc = fitz.open(str(pdf_path))
    for page_num, page in enumerate(doc, 1):
        # Images via PyMuPDF
        for img_idx, img_info in enumerate(page.get_images(full=True)):
            xref = img_info[0]
            base_img = doc.extract_image(xref)
            img_bytes = base_img.get("image")
            if not img_bytes:
                continue
            aid = hash_id(str(pdf_path), str(page_num), str(img_idx), "image")
            img_path = ARTIFACT_DIR / f"{aid}.{base_img.get('ext','png')}"
            img_path.write_bytes(img_bytes)
            artifacts.append({
                "id": aid,
                "type": "image",
                "page": page_num,
                "index": img_idx,
                "content_path": str(img_path),
                "summary_input": f"image_bytes={len(img_bytes)} ext={base_img.get('ext','png')}",
            })
    doc.close()

    if pdfplumber:
        with pdfplumber.open(pdf_path) as pdf:
            for page_num, page in enumerate(pdf.pages, 1):
                for tbl_idx, tbl in enumerate(page.extract_tables() or []):
                    if not tbl or all(not any(c for c in row) for row in tbl):
                        continue
                    aid = hash_id(str(pdf_path), str(page_num), str(tbl_idx), "table")
                    csv_path = ARTIFACT_DIR / f"{aid}.csv"
                    with csv_path.open("w") as fh:
                        for row in tbl:
                            fh.write("\t".join((c or "").replace("\n", " ").replace("\t", " ") for c in row) + "\n")
                    csv_str = csv_path.read_text()[:3000]
                    artifacts.append({
                        "id": aid,
                        "type": "table",
                        "page": page_num,
                        "index": tbl_idx,
                        "content_path": str(csv_path),
                        "summary_input": csv_str,
                    })
    return artifacts


def summarize(artifact: dict) -> str:
    """Summarize an artifact via Claude. Falls back to Heuristik if no API-Key."""
    api_key = os.environ.get("ANTHROPIC_API_KEY")
    if not api_key:
        if artifact["type"] == "table":
            return f"[Heuristic-Summary] Table on page {artifact['page']}, content snippet: {artifact['summary_input'][:200]}"
        return f"[Heuristic-Summary] Image on page {artifact['page']}, no API-Key for Vision"
    try:
        from anthropic import Anthropic
    except ImportError:
        return f"[no-anthropic-sdk] {artifact['type']} on page {artifact['page']}"
    client = Anthropic(api_key=api_key)
    if artifact["type"] == "table":
        prompt = (
            "Fasse die Tabellen-Inhalt in 2-3 deutschen Saetzen zusammen, fokussiert auf: "
            "Header, wichtigste Zahlen, Trends. Maximiere Retrievability fuer spaeteren Vector-Search.\n\n"
            f"Tabelle:\n{artifact['summary_input']}"
        )
        resp = client.messages.create(
            model="claude-haiku-4-5-20251001",
            max_tokens=300,
            messages=[{"role": "user", "content": prompt}],
        )
        return resp.content[0].text
    elif artifact["type"] == "image":
        # Multi-modal: send image bytes
        import base64
        img_b = Path(artifact["content_path"]).read_bytes()
        b64 = base64.b64encode(img_b).decode()
        media_type = "image/png" if artifact["content_path"].endswith(".png") else "image/jpeg"
        resp = client.messages.create(
            model="claude-haiku-4-5-20251001",
            max_tokens=300,
            messages=[{
                "role": "user",
                "content": [
                    {"type": "image", "source": {"type": "base64", "media_type": media_type, "data": b64}},
                    {"type": "text", "text": "Beschreibe das Bild in 2-3 deutschen Saetzen. Fokus: was ist gezeigt, welche Beschriftungen, welche Daten? Maximiere Retrievability fuer Vector-Search."}
                ],
            }],
        )
        return resp.content[0].text
    return "Unsupported artifact type"


def ingest(pdf_path: Path):
    con = ensure_paths()
    print(f"Extracting {pdf_path}...", flush=True)
    artifacts = extract_pdf(pdf_path)
    print(f"  Found {len(artifacts)} artifacts", flush=True)
    for i, a in enumerate(artifacts, 1):
        print(f"  [{i}/{len(artifacts)}] summarizing {a['type']} page {a['page']}", flush=True)
        summary = summarize(a)
        con.execute(
            "INSERT OR REPLACE INTO artifacts (id, source_path, artifact_type, page, artifact_index, artifact_path, summary, created_at) VALUES (?,?,?,?,?,?,?,?)",
            (a["id"], str(pdf_path), a["type"], a["page"], a["index"], a["content_path"], summary, datetime.now(timezone.utc).isoformat()),
        )
    con.commit()
    print(f"OK: {len(artifacts)} artifacts ingested")


def query(question: str, top_n: int = 5):
    con = ensure_paths()
    rows = list(con.execute(
        "SELECT a.id, a.artifact_type, a.page, a.artifact_path, a.summary FROM summaries s JOIN artifacts a ON a.id=s.id WHERE summaries MATCH ? ORDER BY rank LIMIT ?",
        (question, top_n),
    ))
    print(f"Query: {question}")
    print(f"Top-{top_n} matches:")
    for r in rows:
        print(f"  [{r[1]} p{r[2]}] {r[0]} -> {r[3]}")
        print(f"    SUMMARY: {r[4][:200]}...")
    return rows


def main():
    ap = argparse.ArgumentParser()
    sub = ap.add_subparsers(dest="cmd")
    p_ingest = sub.add_parser("ingest")
    p_ingest.add_argument("pdf_path")
    p_query = sub.add_parser("query")
    p_query.add_argument("question")
    p_query.add_argument("--top-n", type=int, default=5)
    p_init = sub.add_parser("init")
    args = ap.parse_args()

    if args.cmd == "init":
        ensure_paths()
        print(f"Initialized DB at {DB_PATH}")
    elif args.cmd == "ingest":
        ingest(Path(args.pdf_path))
    elif args.cmd == "query":
        query(args.question, args.top_n)
    else:
        ap.print_help()
        return 1


if __name__ == "__main__":
    sys.exit(main() or 0)
