"""Analyse d'un fichier importé, exécutée dans un SOUS-PROCESSUS ISOLÉ.

Lancé par le worker avec : environnement vidé, limites mémoire/CPU/taille de fichier, délai maximum,
répertoire de travail dédié, aucun accès à la base ni au réseau applicatif. Entrée : arguments + JSON sur
stdin (indicateurs, options). Sortie : JSON sur stdout. Aucune macro, formule, objet incorporé ni lien
externe n'est exécuté ou suivi : les formules ne sont pas recalculées, seule la valeur en cache est lue.
"""
from __future__ import annotations

import csv
import html
import io
import json
import re
import subprocess
import sys
import warnings
import zipfile
from pathlib import Path
from typing import Any

from . import matching
from .matching import IndicatorRef, Word

MAX_ROWS = 400
MAX_COLS = 30
MAX_PDF_PAGES = 10
DISPLAY_WIDTH = 1600
TESSERACT_TIMEOUT = 90


def _cell(v: Any, t: str | None = None, **flags: Any) -> dict[str, Any]:
    d: dict[str, Any] = {"v": v, "t": t if t is not None else ("" if v is None else str(v))}
    d.update({k: val for k, val in flags.items() if val})
    return d


def _trim(grid: list[list[dict[str, Any]]]) -> list[list[dict[str, Any]]]:
    while grid and all(c.get("t", "") == "" for c in grid[-1]):
        grid.pop()
    width = max((max((i + 1 for i, c in enumerate(r) if c.get("t", "") != ""), default=0) for r in grid), default=0)
    return [r[:width] for r in grid]


# --- Tableurs ------------------------------------------------------------------------------------

def read_csv(path: Path) -> list[dict[str, Any]]:
    raw = path.read_bytes()
    for enc in ("utf-8-sig", "cp1252"):
        try:
            text = raw.decode(enc)
            break
        except UnicodeDecodeError:
            continue
    sample = "\n".join(text.splitlines()[:20])
    delim = max([";", ",", "\t", "|"], key=lambda d: sample.count(d))
    grid = []
    for i, row in enumerate(csv.reader(io.StringIO(text), delimiter=delim)):
        if i >= MAX_ROWS:
            break
        grid.append([_cell(c.strip() if c.strip() else None, c.strip()) for c in row[:MAX_COLS]])
    return [{"kind": "feuille", "index": 0, "name": f"CSV (séparateur « {delim if delim != chr(9) else 'tab'} »)",
             "hidden": False, "grid": _trim(grid), "meta": {"delimiter": delim}}]


def read_xlsx(path: Path) -> list[dict[str, Any]]:
    from openpyxl import load_workbook
    warnings.simplefilter("ignore")
    raw = path.read_bytes()  # le fichier stocké n'a pas d'extension : lecture depuis la mémoire
    wb_f = load_workbook(io.BytesIO(raw), data_only=False, read_only=False, keep_links=False, keep_vba=False)
    wb_v = load_workbook(io.BytesIO(raw), data_only=True, read_only=False, keep_links=False, keep_vba=False)
    sources = []
    for idx, ws in enumerate(wb_f.worksheets):
        wv = wb_v.worksheets[idx]
        merged_tl: dict[tuple[int, int], bool] = {}
        merged_any: set[tuple[int, int]] = set()
        for rng in ws.merged_cells.ranges:
            for r in range(rng.min_row, min(rng.max_row, MAX_ROWS) + 1):
                for c in range(rng.min_col, min(rng.max_col, MAX_COLS) + 1):
                    merged_any.add((r, c))
            merged_tl[(rng.min_row, rng.min_col)] = True
        grid = []
        formulas = 0
        max_row = min(ws.max_row or 0, MAX_ROWS)
        max_col = min(ws.max_column or 0, MAX_COLS)
        for r in range(1, max_row + 1):
            hidden_row = bool(ws.row_dimensions[r].hidden) if r in ws.row_dimensions else False
            row = []
            for c in range(1, max_col + 1):
                fv = ws.cell(row=r, column=c).value
                vv = wv.cell(row=r, column=c).value
                formula = fv if isinstance(fv, str) and fv.startswith("=") else None
                if formula:
                    formulas += 1
                value = vv if formula else fv
                if hasattr(value, "isoformat"):
                    text = value.strftime("%d/%m/%Y") if hasattr(value, "strftime") else str(value)
                    value = text
                if isinstance(value, str):
                    value = value.strip() or None
                row.append(_cell(value, None if value is not None else "", f=formula,
                                 merged=(r, c) in merged_any, hidden_row=hidden_row))
            grid.append(row)
        sources.append({
            "kind": "feuille", "index": idx, "name": ws.title, "hidden": ws.sheet_state != "visible",
            "grid": _trim(grid),
            "meta": {"formulas": formulas, "merged_ranges": len(ws.merged_cells.ranges),
                     "truncated": (ws.max_row or 0) > MAX_ROWS or (ws.max_column or 0) > MAX_COLS,
                     "external_links": bool(getattr(wb_f, "_external_links", []))},
        })
    return sources


ODS = {
    "table": "urn:oasis:names:tc:opendocument:xmlns:table:1.0",
    "office": "urn:oasis:names:tc:opendocument:xmlns:office:1.0",
    "text": "urn:oasis:names:tc:opendocument:xmlns:text:1.0",
    "style": "urn:oasis:names:tc:opendocument:xmlns:style:1.0",
}


def _q(ns: str, tag: str) -> str:
    return f"{{{ODS[ns]}}}{tag}"


def read_ods(path: Path) -> list[dict[str, Any]]:
    from defusedxml import ElementTree as SafeET
    with zipfile.ZipFile(path) as zf:
        root = SafeET.fromstring(zf.read("content.xml"), forbid_dtd=True)
    hidden_styles = set()
    for st in root.iter(_q("style", "style")):
        props = st.find(_q("style", "table-properties"))
        if props is not None and props.get(_q("table", "display")) == "false":
            hidden_styles.add(st.get(_q("style", "name")))
    sources = []
    for idx, table in enumerate(root.iter(_q("table", "table"))):
        grid: list[list[dict[str, Any]]] = []
        formulas = merged = 0
        empty_streak = 0
        for row in table.iter(_q("table", "table-row")):
            if len(grid) >= MAX_ROWS or empty_streak > 60:
                break
            rep = min(int(row.get(_q("table", "number-rows-repeated"), "1")), MAX_ROWS)
            hidden_row = row.get(_q("table", "visibility")) in {"collapse", "filter"}
            cells: list[dict[str, Any]] = []
            for cell in row:
                if len(cells) >= MAX_COLS:
                    break
                tag = cell.tag.split("}")[-1]
                if tag not in {"table-cell", "covered-table-cell"}:
                    continue
                crep = min(int(cell.get(_q("table", "number-columns-repeated"), "1")), MAX_COLS)
                vtype = cell.get(_q("office", "value-type"))
                text = "\n".join("".join(p.itertext()) for p in cell.findall(_q("text", "p"))).strip()
                value: Any = None
                if vtype in {"float", "percentage", "currency"}:
                    try:
                        value = float(cell.get(_q("office", "value"), "nan"))
                        value = int(value) if value == int(value) else value
                    except ValueError:
                        value = text or None
                elif text:
                    value = text
                formula = cell.get(_q("table", "formula"))
                if formula:
                    formulas += 1
                is_merged = tag == "covered-table-cell" or int(cell.get(_q("table", "number-columns-spanned"), "1")) > 1 \
                    or int(cell.get(_q("table", "number-rows-spanned"), "1")) > 1
                merged += int(is_merged)
                for _ in range(crep):
                    if len(cells) >= MAX_COLS:
                        break
                    cells.append(_cell(value, text, f=formula, merged=is_merged, hidden_row=hidden_row))
            empty = all(c["t"] == "" for c in cells)
            empty_streak = empty_streak + rep if empty else 0
            for _ in range(rep if not empty else min(rep, 3)):
                if len(grid) >= MAX_ROWS:
                    break
                grid.append([dict(c) for c in cells])
        sources.append({
            "kind": "feuille", "index": idx, "name": table.get(_q("table", "name"), f"Feuille {idx + 1}"),
            "hidden": table.get(_q("table", "style-name")) in hidden_styles, "grid": _trim(grid),
            "meta": {"formulas": formulas, "merged_cells": merged},
        })
    return sources


# --- Images et OCR -------------------------------------------------------------------------------

def _tesseract(binary: str, image: Path, args: list[str]) -> str:
    out = subprocess.run([binary, str(image), "stdout", *args], capture_output=True, timeout=TESSERACT_TIMEOUT,
                         check=False)
    return out.stdout.decode("utf-8", errors="replace")


def detect_rotation(binary: str, image: Path) -> tuple[int, float]:
    text = _tesseract(binary, image, ["--psm", "0", "-l", "osd"])
    rot = re.search(r"Rotate:\s*(\d+)", text)
    conf = re.search(r"Orientation confidence:\s*([\d.]+)", text)
    return (int(rot.group(1)) if rot else 0), (float(conf.group(1)) if conf else 0.0)


def ocr_words(binary: str, image: Path, scale: float, psm: int) -> list[Word]:
    tsv = _tesseract(binary, image, ["-l", "fra", "--psm", str(psm), "--oem", "1", "tsv"])
    words = []
    for line in tsv.splitlines()[1:]:
        parts = line.split("\t")
        if len(parts) < 12 or parts[0] != "5":
            continue
        text = parts[11].strip()
        conf = float(parts[10]) if parts[10] not in {"", "-1"} else -1
        if not text or conf < 0:
            continue
        left, top, w, h = (int(x) for x in parts[6:10])
        words.append(Word(text, left / scale, top / scale, (left + w) / scale, (top + h) / scale, conf / 100))
    return words


def prepare_image(src: Path, workdir: Path, name: str, rotation: int | None, contrast: bool, binary: str
                  ) -> tuple[dict[str, Any], Path, float, list[str]]:
    """Réencode l'image (élimine métadonnées et contenus parasites), l'oriente et prépare la version OCR."""
    from PIL import Image, ImageEnhance, ImageFilter, ImageOps
    Image.MAX_IMAGE_PIXELS = 40_000_000
    notes: list[str] = []
    with Image.open(src) as im:
        im.verify()
    with Image.open(src) as im:
        im = ImageOps.exif_transpose(im).convert("RGB")
    if max(im.size) > 4000:
        im.thumbnail((4000, 4000))
    probe = workdir / f"{name}-probe.png"
    im.save(probe, "PNG")
    applied = 0
    if rotation is None:
        rot, conf = detect_rotation(binary, probe)
        if rot and conf >= 1.5:
            applied = rot
            notes.append(f"Orientation corrigée automatiquement ({rot}°) : utilisez « Pivoter » si besoin.")
    else:
        applied = rotation % 360
    if applied:
        im = im.rotate(-applied, expand=True)
    probe.unlink(missing_ok=True)
    display = im.copy()
    if display.width > DISPLAY_WIDTH:
        display.thumbnail((DISPLAY_WIDTH, DISPLAY_WIDTH * 4))
    display_name = f"{name}.png"
    display.save(workdir / display_name, "PNG", optimize=True)
    gray = ImageOps.autocontrast(ImageOps.grayscale(im), cutoff=1)
    if contrast:
        gray = ImageEnhance.Contrast(gray).enhance(1.8).filter(ImageFilter.SHARPEN)
        notes.append("Contraste renforcé pour l'OCR.")
    target = 2400
    factor = max(0.5, min(3.0, target / gray.width))
    if abs(factor - 1) > 0.05:
        gray = gray.resize((int(gray.width * factor), int(gray.height * factor)), Image.LANCZOS)
    ocr_path = workdir / f"{name}-ocr.png"
    gray.save(ocr_path, "PNG")
    to_display = display.width / im.width
    info = {"width": display.width, "height": display.height, "image": display_name, "rotation": applied}
    return info, ocr_path, factor / to_display, notes


def ocr_page(binary: str, ocr_path: Path, scale: float, indicators: list[IndicatorRef], source_index: int
             ) -> tuple[list[dict[str, Any]], list[str], list[str], list[Word]]:
    best: tuple[list[dict[str, Any]], list[str], list[str], list[Word]] | None = None
    for psm in (6, 4):
        words = ocr_words(binary, ocr_path, scale, psm)
        props, warns, unused = matching.proposals_from_words(words, indicators, source_index, ocr=True)
        recognized = sum(1 for p in props if p.get("indicator_code"))
        if best is None or recognized > sum(1 for p in best[0] if p.get("indicator_code")):
            best = (props, warns, unused, words)
        if recognized >= 10:
            break
    assert best is not None
    reread_numbers(binary, ocr_path, scale, best[0])
    ocr_path.unlink(missing_ok=True)
    return best


def reread_numbers(binary: str, ocr_path: Path, scale: float, proposals: list[dict[str, Any]]) -> None:
    """Seconde lecture ciblée de la zone « Nombre » (chiffres uniquement) quand la première est douteuse.

    Le résultat reste une PROPOSITION signalée : il n'est jamais accepté sans confirmation humaine."""
    from PIL import Image
    with Image.open(ocr_path) as img:
        img = img.copy()
    for k, p in enumerate(proposals):
        zone, box = p.get("value_zone"), p.get("bbox")
        if not p.get("indicator_code") or not zone or not box:
            continue
        if p.get("value") is not None and not p.get("value_error") and p.get("confidence", 0) >= 0.8:
            continue
        cy = box["y"] + box["h"] / 2
        half = max(box["h"] / 2, 18 / scale)
        crop = img.crop((int(zone[0] * scale), int((cy - half) * scale), int(zone[1] * scale),
                         int((cy + half) * scale)))
        if crop.width < 5 or crop.height < 5:
            continue
        cpath = ocr_path.with_name(f"cell-{k}.png")
        crop.save(cpath)
        tsv = _tesseract(binary, cpath, ["--psm", "7", "-c", "tessedit_char_whitelist=0123456789", "tsv"])
        cpath.unlink(missing_ok=True)
        digits, confs = "", []
        for line in tsv.splitlines()[1:]:
            parts = line.split("\t")
            if len(parts) >= 12 and parts[0] == "5" and parts[11].strip():
                digits += parts[11].strip()
                confs.append(float(parts[10]) / 100)
        if not digits.isdigit():
            continue
        conf = min(confs) if confs else 0.0
        before = p.get("raw_value") or "(rien)"
        p["warnings"] = [w for w in p["warnings"] if not w.startswith(("Aucun nombre lu", "Caractère ambigu",
                                                                       "Chiffre lu avec"))]
        p["warnings"].append(f"Valeur relue par une seconde lecture ciblée (lu d'abord « {before} ») : "
                             "vérifiez sur l'image.")
        p.update(raw_value=digits, value=int(digits), value_error=None,
                 confidence=round(min(p.get("match_score", 0), conf, 0.85), 3),
                 value_bbox={"x": zone[0], "y": round(cy - half, 1), "w": round(zone[1] - zone[0], 1),
                             "h": round(2 * half, 1)})


# --- PDF -----------------------------------------------------------------------------------------

def _run(args: list[str], timeout: int = 60) -> subprocess.CompletedProcess[bytes]:
    return subprocess.run(args, capture_output=True, timeout=timeout, check=False)


def pdf_pages(path: Path) -> int:
    out = _run(["pdfinfo", str(path)], 30).stdout.decode("utf-8", errors="replace")
    m = re.search(r"^Pages:\s+(\d+)", out, re.M)
    if not m:
        raise ValueError("PDF illisible.")
    return int(m.group(1))


WORD_RE = re.compile(r'<word xMin="([\d.]+)" yMin="([\d.]+)" xMax="([\d.]+)" yMax="([\d.]+)">(.*?)</word>', re.S)
PAGE_RE = re.compile(r'<page width="([\d.]+)" height="([\d.]+)">(.*?)</page>', re.S)


def pdf_text_words(path: Path, workdir: Path, pages: int) -> list[tuple[float, float, list[Word]]]:
    out = workdir / "text.html"
    _run(["pdftotext", "-bbox-layout", "-f", "1", "-l", str(pages), str(path), str(out)], 60)
    if not out.is_file():
        return []
    content = out.read_text(encoding="utf-8", errors="replace")
    out.unlink()
    result = []
    for m in PAGE_RE.finditer(content):
        words = [Word(html.unescape(w[4]), float(w[0]), float(w[1]), float(w[2]), float(w[3]), 1.0)
                 for w in WORD_RE.findall(m.group(3))]
        result.append((float(m.group(1)), float(m.group(2)), words))
    return result


# --- Point d'entrée ------------------------------------------------------------------------------

def analyse(path: Path, fmt: str, workdir: Path, payload: dict[str, Any]) -> dict[str, Any]:
    indicators = [IndicatorRef(**i) for i in payload["indicators"]]
    opts = payload.get("options", {})
    binary = opts.get("tesseract", "tesseract")
    result: dict[str, Any] = {"format": fmt, "sources": [], "proposals": [], "warnings": [], "unused_lines": []}
    texts_for_week: list[tuple[str, str]] = [(opts.get("original_name", ""), "nom du fichier")]
    if fmt in {"csv", "xlsx", "ods"}:
        sources = {"csv": read_csv, "xlsx": read_xlsx, "ods": read_ods}[fmt](path)
        for src in sources:
            props, warns, cols = matching.proposals_from_grid(src["grid"], indicators, src["index"])
            src["meta"]["columns"] = cols
            src["meta"]["warnings"] = warns
            result["proposals"].extend(props)
            if src["hidden"]:
                src["meta"]["warnings"].append("Feuille masquée dans le classeur.")
            for r, row in enumerate(src["grid"][:15]):
                for c, cell in enumerate(row[:10]):
                    if cell.get("t"):
                        texts_for_week.append((cell["t"], f"{src['name']}!{matching.col_letter(c)}{r + 1}"))
        result["sources"] = sources
    elif fmt in {"png", "jpeg"}:
        info, ocr_path, scale, notes = prepare_image(path, workdir, "image-1", opts.get("rotation"),
                                                     bool(opts.get("contrast")), binary)
        props, warns, unused, words = ocr_page(binary, ocr_path, scale, indicators, 0)
        result["sources"] = [{"kind": "image", "index": 0, "name": "Image", "hidden": False, **info,
                              "meta": {"warnings": notes + warns, "words": len(words), "ocr": True,
                                       "mean_confidence": round(sum(w.conf for w in words) / len(words), 3) if words else 0}}]
        result["proposals"] = props
        result["unused_lines"] = unused
        texts_for_week += [(" ".join(w.text for w in row), "image") for row in matching.group_rows(words)[:15]]
    elif fmt == "pdf":
        pages = pdf_pages(path)
        if pages > MAX_PDF_PAGES:
            result["warnings"].append(f"Seules les {MAX_PDF_PAGES} premières pages sont analysées.")
            pages = MAX_PDF_PAGES
        text_pages = pdf_text_words(path, workdir, pages)
        for p in range(1, pages + 1):
            _run(["pdftoppm", "-f", str(p), "-l", str(p), "-r", "110", "-png", "-singlefile", str(path),
                  str(workdir / f"page-{p}")], 60)
            display = workdir / f"page-{p}.png"
            from PIL import Image
            Image.MAX_IMAGE_PIXELS = 40_000_000
            with Image.open(display) as im:
                width, height = im.size
            words: list[Word] = []
            ocr_used = False
            notes: list[str] = []
            if p - 1 < len(text_pages) and len(text_pages[p - 1][2]) >= 8:
                pw, _ph, pwords = text_pages[p - 1]
                k = width / pw
                words = [Word(w.text, w.x0 * k, w.y0 * k, w.x1 * k, w.y1 * k, 1.0) for w in pwords]
                props, warns, unused = matching.proposals_from_words(words, indicators, p - 1, ocr=False)
                notes.append("Texte extrait directement du PDF (pas d'OCR).")
            else:
                ocr_used = True
                _run(["pdftoppm", "-f", str(p), "-l", str(p), "-r", "250", "-gray", "-png", "-singlefile",
                      str(path), str(workdir / f"ocr-{p}")], 90)
                ocr_path = workdir / f"ocr-{p}.png"
                props, warns, unused, words = ocr_page(binary, ocr_path, 250 / 110, indicators, p - 1)
                notes.append("Page numérisée : texte reconnu par OCR local.")
            result["sources"].append({
                "kind": "page", "index": p - 1, "name": f"Page {p}", "hidden": False, "width": width,
                "height": height, "image": f"page-{p}.png",
                "meta": {"warnings": notes + warns, "words": len(words), "ocr": ocr_used,
                         "mean_confidence": round(sum(w.conf for w in words) / len(words), 3) if words else 0}})
            result["proposals"].extend(props)
            result["unused_lines"].extend(unused)
            texts_for_week += [(" ".join(w.text for w in row), f"page {p}") for row in matching.group_rows(words)[:15]]
    else:
        raise ValueError("Format non pris en charge.")
    result["week"] = matching.detect_week(texts_for_week)
    # Sélection de la source la plus pertinente ; les autres sont ignorées par défaut (modifiable).
    counts: dict[int, int] = {}
    for p in result["proposals"]:
        if p.get("indicator_code"):
            counts[p["source_index"]] = counts.get(p["source_index"], 0) + 1
    selected = max(counts, key=lambda k: counts[k]) if counts else 0
    result["selected_source"] = selected
    for p in result["proposals"]:
        if p["source_index"] != selected:
            p["decision"] = "ignore"
    return result


def main() -> None:
    path, fmt, workdir = Path(sys.argv[1]), sys.argv[2], Path(sys.argv[3])
    payload = json.loads(sys.stdin.read() or "{}")
    try:
        result = analyse(path, fmt, workdir, payload)
        sys.stdout.write(json.dumps({"ok": True, "result": result}, ensure_ascii=False, default=str))
    except MemoryError:
        sys.stdout.write(json.dumps({"ok": False, "error": "Fichier trop complexe : limite mémoire atteinte."}))
    except subprocess.TimeoutExpired:
        sys.stdout.write(json.dumps({"ok": False, "error": "Analyse trop longue : délai maximum dépassé."}))
    except Exception as e:  # message générique, pas de trace côté utilisateur
        sys.stderr.write(f"{type(e).__name__}: {e}\n")
        sys.stdout.write(json.dumps({"ok": False, "error": "Le fichier n'a pas pu être analysé (contenu illisible "
                                                          "ou structure non reconnue)."}))


if __name__ == "__main__":
    main()
