#!/usr/bin/env python3
"""
OCR de grilla + QR de una factura/boleta, para "grounding" de la extraccion por vision.

Uso: python invoice_ocr.py <ruta_imagen>
Salida (stdout, JSON):
  {
    "qr": "<texto del QR o null>",
    "rows": [["token", ...], ...],
    "spatial_rows": [{"y": 123, "tokens": [{"text":"...", "x":45, "confidence":0.9}]}]
  }

RapidOCR (ONNX, CPU) lee el texto POR POSICION, asi que preserva la alineacion de columnas
que el LLM a veces confunde. Le pasamos estas filas al LLM como referencia de los numeros
exactos. Best-effort: si falta rapidocr/cv2, devuelve lo que pueda.
"""
import sys
import json


def read_qr(path):
    try:
        import cv2
    except Exception:
        return None
    img = cv2.imread(path)
    if img is None:
        return None
    det = cv2.QRCodeDetector()
    try:
        data, _, _ = det.detectAndDecode(img)
        if data:
            return data
    except Exception:
        pass
    try:
        ok, infos, _, _ = det.detectAndDecodeMulti(img)
        if ok:
            for s in infos:
                if s:
                    return s
    except Exception:
        pass
    try:
        g = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
        for scale in (1.5, 2.0, 0.7):
            gg = cv2.resize(g, None, fx=scale, fy=scale)
            d2, _, _ = det.detectAndDecode(gg)
            if d2:
                return d2
    except Exception:
        pass
    return None


def group_tokens(toks):
    if not toks:
        return {"rows": [], "spatial_rows": []}

    # Agrupa contra una línea base fija, no contra el promedio móvil. El promedio
    # móvil podía encadenar gradualmente dos renglones contiguos en una sola fila.
    heights = sorted(t["h"] for t in toks)
    med = heights[len(heights) // 2] or 10
    th = max(5.0, med * 0.48)

    toks.sort(key=lambda z: z["y"])
    rows = []
    cur = []
    for t in toks:
        if not cur or abs(t["y"] - cur_y) <= th:
            cur.append(t)
            cur_y = sorted(x["y"] for x in cur)[len(cur) // 2]
        else:
            rows.append(cur)
            cur = [t]
            cur_y = t["y"]
    if cur:
        rows.append(cur)

    rows = [sorted(r, key=lambda z: z["x"]) for r in rows]
    return {
        "rows": [[x["t"] for x in r] for r in rows],
        "spatial_rows": [
            {
                "y": round(sum(x["y"] for x in r) / len(r), 1),
                "tokens": [
                    {"text": x["t"], "x": round(x["x"], 1), "confidence": round(x["conf"], 3)}
                    for x in r
                ],
            }
            for r in rows
        ],
    }


def ocr_rows(path):
    try:
        from rapidocr_onnxruntime import RapidOCR
    except Exception:
        return {"rows": [], "spatial_rows": []}
    ocr = RapidOCR()
    res, _ = ocr(path)
    if not res:
        return {"rows": [], "spatial_rows": []}

    toks = []
    for box, txt, conf in res:
        xs = [p[0] for p in box]
        ys = [p[1] for p in box]
        toks.append({
            "t": txt,
            "x": sum(xs) / 4.0,
            "y": sum(ys) / 4.0,
            "h": max(ys) - min(ys),
            "conf": float(conf),
        })

    return group_tokens(toks)


def main():
    try:
        sys.stdout.reconfigure(encoding="utf-8", errors="replace")
    except Exception:
        pass
    out = {"qr": None, "rows": []}
    if len(sys.argv) < 2:
        print(json.dumps(out))
        return
    path = sys.argv[1]
    try:
        out["qr"] = read_qr(path)
    except Exception:
        pass
    try:
        detected = ocr_rows(path)
        out["rows"] = detected.get("rows", [])
        out["spatial_rows"] = detected.get("spatial_rows", [])
    except Exception as e:
        out["ocr_error"] = str(e)[:150]
    print(json.dumps(out, ensure_ascii=False))


if __name__ == "__main__":
    main()
