# ai_google_vision.py
import os, json, textwrap
from pathlib import Path
from typing import List, Dict
import fitz  # PyMuPDF
from PIL import Image  # noqa: F401 (optional)
import google.generativeai as genai

# Bevorzugt: Gemini 2.5 Flash (Fallback: 1.5 Flash)
PREFERRED_MODEL = "models/gemini-2.5-flash"
FALLBACK_MODEL  = "models/gemini-1.5-flash"

def _ensure_key_and_model():
    key = os.environ.get("GOOGLE_API_KEY")
    if not key:
        raise RuntimeError("GOOGLE_API_KEY ist nicht gesetzt (.env / setx).")
    genai.configure(api_key=key)
    try:
        return genai.GenerativeModel(PREFERRED_MODEL)
    except Exception:
        return genai.GenerativeModel(FALLBACK_MODEL)

def _pdf_to_png_bytes(pdf_path: Path, dpi: int = 220, max_pages: int | None = None) -> List[bytes]:
    pdf_path = Path(pdf_path)
    doc = fitz.open(str(pdf_path))
    images: List[bytes] = []
    zoom = dpi / 72.0
    mat = fitz.Matrix(zoom, zoom)
    pages = range(len(doc)) if max_pages is None else range(min(max_pages, len(doc)))
    for i in pages:
        page = doc.load_page(i)
        pix = page.get_pixmap(matrix=mat, alpha=False)
        images.append(pix.tobytes("png"))
    doc.close()
    return images

# -------- Offerten: Positionen mit EP/GP/Text (OCR + Parsing) -----------------
def extract_lv_from_pdf(pdf_path: Path, dpi: int = 220, batch_size: int = 6) -> Dict[str, Dict]:
    """
    Liest gescannte Offerten (DE) und liefert:
      { pos: {'EP': float|None, 'GP': float|None, 'text': str} }
    """
    model = _ensure_key_and_model()
    imgs = _pdf_to_png_bytes(pdf_path, dpi=dpi)

    sys_prompt = textwrap.dedent("""
        Du bist ein deutschsprachiger Experte für NPK-Leistungsverzeichnisse und Offerten.
        Aufgabe: Erkenne aus eingescannten Offerten mit ggf. handschriftlichen Preisen die NPK-Positionen.
        Antworte ausschließlich auf DEUTSCH und ausschließlich als JSON.

        Für jede Position:
        - "pos"  : reine Positionsnummer (nur Ziffern, keine Punkte/Leerzeichen)
        - "text" : kurzer Positionstext (max. 200 Zeichen)
        - "EP"   : Einheitspreis (float, CHF ohne Tausendertrennzeichen) oder null
        - "GP"   : Gesamtpreis/Positionssumme (float) oder null

        JSON (genau so, ohne Kommentare):
        {"positions":[{"pos":"...","text":"...","EP":<float|null>,"GP":<float|null>}, ...]}
    """).strip()

    all_positions: list[dict] = []
    for i in range(0, len(imgs), batch_size):
        parts = [sys_prompt]
        for b in imgs[i:i+batch_size]:
            parts.append({"mime_type": "image/png", "data": b})
        resp = model.generate_content(parts, generation_config={"response_mime_type": "application/json"})
        try:
            data = json.loads(resp.text or "{}")
            ps = data.get("positions", [])
            if isinstance(ps, list):
                all_positions.extend(ps)
        except Exception:
            continue

    result: Dict[str, Dict] = {}
    for p in all_positions:
        pos = str(p.get("pos", "")).strip()
        if not pos.isdigit():
            continue
        entry = result.setdefault(pos, {"EP": None, "GP": None, "text": ""})
        txt = (p.get("text") or "").strip()
        if txt and not entry["text"]:
            entry["text"] = txt[:200]
        ep = p.get("EP"); gp = p.get("GP")
        if isinstance(ep, (int, float)): entry["EP"] = float(ep)
        if isinstance(gp, (int, float)): entry["GP"] = float(gp)
    return result

# -------- AUSSCHREIBUNG: Positionen mit Menge/Einheit (OCR + Parsing) --------
def extract_lv_qty_from_pdf(pdf_path: Path, dpi: int = 220, batch_size: int = 6) -> Dict[str, Dict]:
    """
    Liest aus der AUSSCHREIBUNG alle Positionen, die eine ausgeschriebene Menge besitzen:
      { pos: {'qty': float, 'unit': str, 'text': str} }
    """
    model = _ensure_key_and_model()
    imgs = _pdf_to_png_bytes(pdf_path, dpi=dpi)

    sys_prompt = textwrap.dedent("""
        Du bist ein deutschsprachiger Experte für NPK-Leistungsverzeichnisse.
        Aufgabe: Lies aus einer AUSSCHREIBUNG alle Positionen, die eine ausgeschriebene Menge besitzen.
        Antworte ausschließlich auf DEUTSCH und ausschließlich als JSON.

        Für jede Position erfasse:
        - "pos"   : reine Positionsnummer (nur Ziffern, keine Punkte/Leerzeichen)
        - "text"  : kurzer Positionstext (max. 200 Zeichen)
        - "qty"   : ausgeschriebene Menge als float
        - "unit"  : Mengeneinheit (z. B. m, m2, Stk, kg)

        JSON (ohne Kommentare):
        {"positions":[{"pos":"...","text":"...","qty":<float>,"unit":"..."}, ...]}

        NUR Positionen mit Menge liefern.
    """).strip()

    all_positions: list[dict] = []
    for i in range(0, len(imgs), batch_size):
        parts = [sys_prompt]
        for b in imgs[i:i+batch_size]:
            parts.append({"mime_type": "image/png", "data": b})
        resp = model.generate_content(parts, generation_config={"response_mime_type": "application/json"})
        try:
            data = json.loads(resp.text or "{}")
            ps = data.get("positions", [])
            if isinstance(ps, list):
                all_positions.extend(ps)
        except Exception:
            continue

    result: Dict[str, Dict] = {}
    for p in all_positions:
        pos = str(p.get("pos", "")).strip()
        if not pos.isdigit():
            continue
        try:
            qty = float(p.get("qty"))
        except Exception:
            continue  # nur mit gültiger Menge
        unit = str(p.get("unit") or "").strip()
        txt  = (p.get("text") or "").strip()
        result[pos] = {"qty": qty, "unit": unit, "text": txt[:200]}
    return result

# -------- Kurze DE-Begründungen für Ausreißer --------------------------------
def explain_outliers(offer_data: Dict[str, Dict], detail_rows: List[Dict]) -> Dict[str, str]:
    """
    Kurze Begründungen (DE) für Ausreißer je Position.
    Return: { pos: "1–2 Sätze" }
    """
    model = _ensure_key_and_model()

    payload = []
    for row in detail_rows:
        if not row.get("outliers"):
            continue
        pos = row["pos"]
        entry = {"pos": pos, "median": row["median"], "vendors":[]}
        for vendor, posmap in offer_data.items():
            ep = (posmap.get(pos) or {}).get("EP")
            txt= (posmap.get(pos) or {}).get("text")
            if isinstance(ep, (int, float)):
                entry["vendors"].append({"name": vendor, "EP": ep, "text": (txt or "")[:200]})
        if entry["vendors"]:
            payload.append(entry)
    if not payload:
        return {}

    sys_prompt = (
        "Du bist ein deutschsprachiger Kalkulationsexperte. "
        "Analysiere die Preis-Ausreißer je Position und gib pro Position eine sehr kurze, plausible Begründung (1–2 Sätze) "
        "in DEUTSCH, z. B. Leistungsumfang, Qualitätsniveau, Nebenleistungen, Ausführungsdetails, Mengenannahmen. "
        "Antworte ausschließlich als JSON im Format: "
        '{"reasons":[{"pos":"<pos>","reason":"<kurze deutsche Begründung>"}]}'
    )
    user = json.dumps({"positions": payload}, ensure_ascii=False)
    resp = model.generate_content([sys_prompt, user], generation_config={"response_mime_type":"application/json"})
    out = {}
    try:
        data = json.loads(resp.text or "{}").get("reasons", [])
        for item in data:
            out[str(item.get("pos"))] = str(item.get("reason","")).strip()
    except Exception:
        pass
    return out
