#!/usr/bin/env python3
"""Regressione del riconoscimento: rifa' in Python ESATTAMENTE cio' che fa il worker.

Serve a rispondere a una sola domanda: dopo aver toccato la galleria, la stessa foto
da' ancora le stesse carte? (la prova su telefono non e' ripetibile da qui).

Pipeline replicata da pwa/app.js captureSquare() + pwa/recognize.js defaultGrid()
+ pwa/recognize.worker.js run(): ritaglio quadrato centrale -> 1600x1600 -> grigio,
celle della griglia -> larghezza 400 -> ORB 400 keypoint -> knnMatch/Lowe 0.75
contro le sole carte dello stesso tipo.

  /usr/bin/python3 recognition/test_gallery.py assets/photos/tableaus/tavolo.jpg
  /usr/bin/python3 recognition/test_gallery.py FOTO --gallery /path/vecchia   # confronto
  /usr/bin/python3 recognition/test_gallery.py FOTO --expect 68,24,28,23,42,25,7,2
"""
import os, sys, json, argparse, cv2, numpy as np

ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
RATIO, NFEAT_Q, SZ = 0.75, 400, 1600


def default_grid():
    """copia 1:1 di REC.defaultGrid() — 6 Santuari in alto, 8 Regioni 4x2."""
    M, GAP, SAN_RATIO = 0.020, 0.012, 68 / 44
    cells = []

    def cell(t, gx, gy, gw, gh, px, py):
        cells.append({"type": t, "x": gx - px, "y": gy - py, "w": gw + 2 * px, "h": gh + 2 * py})
    sw = (1 - 2 * M - 5 * GAP) / 6
    for i in range(6):
        cell("sanctuary", M + i * (sw + GAP), 0.105, sw, sw * SAN_RATIO, 0.005, 0.018)
    rw, ry, rgap = (1 - 2 * M - 3 * GAP) / 4, 0.418, 0.042
    for r in range(2):
        for c in range(4):
            cell("region", M + c * (rw + GAP), ry + r * (rw + rgap), rw, rw, 0.005, 0.012)
    return cells


def square(path):
    im = cv2.imread(path)
    if im is None:
        raise SystemExit("foto illeggibile: " + path)
    h, w = im.shape[:2]
    s = min(w, h)
    im = im[(h - s) // 2:(h - s) // 2 + s, (w - s) // 2:(w - s) // 2 + s]
    return cv2.cvtColor(cv2.resize(im, (SZ, SZ), interpolation=cv2.INTER_AREA), cv2.COLOR_BGR2GRAY)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("foto")
    ap.add_argument("--gallery", default=os.path.join(ROOT, "pwa", "data"),
                    help="cartella con orb-gallery.json/.bin (default: quella della PWA)")
    ap.add_argument("--expect", help="numeri Regione attesi, separati da virgola")
    a = ap.parse_args()

    meta = json.load(open(os.path.join(a.gallery, "orb-gallery.json")))
    buf = np.fromfile(os.path.join(a.gallery, "orb-gallery.bin"), dtype=np.uint8)
    W = meta["orbW"]
    gal = [(c["id"], c["type"], buf[c["off"]:c["off"] + c["n"] * 32].reshape(c["n"], 32))
           for c in meta["cards"] if c["n"] > 0]
    print("galleria %s: %d carte" % (a.gallery, len(gal)))

    gray = square(a.foto)
    orb = cv2.ORB_create(NFEAT_Q)
    bf = cv2.BFMatcher(cv2.NORM_HAMMING, crossCheck=False)
    regioni = []
    for cel in default_grid():
        ih, iw = gray.shape
        x, y = max(0, min(iw - 1, round(cel["x"] * iw))), max(0, min(ih - 1, round(cel["y"] * ih)))
        w, h = max(1, min(iw - x, round(cel["w"] * iw))), max(1, min(ih - y, round(cel["h"] * ih)))
        roi = cv2.resize(gray[y:y + h, x:x + w], (W, max(1, round(W * h / w))), interpolation=cv2.INTER_AREA)
        _, q = orb.detectAndCompute(roi, None)
        best, bn, second = None, 0, 0
        if q is not None and len(q):
            for cid, tipo, g in gal:
                if tipo != cel["type"]:
                    continue
                good = sum(1 for m in bf.knnMatch(q, g, k=2)
                           if len(m) >= 2 and m[0].distance < RATIO * m[1].distance)
                if good > bn:
                    bn, second, best = good, bn, cid
                elif good > second:
                    second = good
        print("  %-9s %-5s score %3d  margine %3d" % (cel["type"], best, bn, bn - second))
        if cel["type"] == "region":
            regioni.append(best)

    if a.expect:
        atteso = [s.strip() for s in a.expect.split(",")]
        ok = regioni == atteso
        print("\nRegioni: %s\nAtteso : %s\n%s" % (regioni, atteso, "OK" if ok else "REGRESSIONE"))
        sys.exit(0 if ok else 1)
    print("\nRegioni:", regioni)


if __name__ == "__main__":
    main()
