#!/usr/bin/env python3
"""Riconoscimento Santuari FARAWAY via embedding DINOv2 + nearest-neighbor.
Uso:
  python dino_match.py --selftest          # collaudo: galleria vs versioni aumentate
  python dino_match.py --gallery DIR --query IMG   # match di una carta contro la galleria
La galleria e una cartella di immagini di riferimento (nome file = id carta).
"""
import sys, os, glob, argparse, numpy as np
from PIL import Image, ImageEnhance
import torch
from transformers import AutoImageProcessor, AutoModel

MODEL = "facebook/dinov2-small"
_proc = _model = None
def _load():
    global _proc, _model
    if _model is None:
        _proc = AutoImageProcessor.from_pretrained(MODEL)
        _model = AutoModel.from_pretrained(MODEL).eval()
    return _proc, _model

@torch.no_grad()
def embed(img: Image.Image) -> np.ndarray:
    proc, model = _load()
    inp = proc(images=img.convert("RGB"), return_tensors="pt")
    out = model(**inp)
    v = out.last_hidden_state[:, 0]          # CLS token = embedding immagine
    v = torch.nn.functional.normalize(v, dim=-1)
    return v[0].cpu().numpy()

def build_gallery(folder):
    g = {}
    for p in sorted(glob.glob(os.path.join(folder, "*"))):
        if p.lower().endswith((".jpg", ".jpeg", ".png")):
            g[os.path.splitext(os.path.basename(p))[0]] = embed(Image.open(p))
    return g

def match(q_emb, gallery, topk=3):
    sims = sorted(((float(np.dot(q_emb, e)), name) for name, e in gallery.items()), reverse=True)
    return sims[:topk]

def augment(img, k):
    """Piccole variazioni per simulare scatti diversi."""
    import random
    random.seed(k)
    im = img.convert("RGB")
    im = im.rotate(random.uniform(-12, 12), expand=False, fillcolor=(20, 60, 30))
    im = ImageEnhance.Brightness(im).enhance(random.uniform(0.75, 1.25))
    im = ImageEnhance.Contrast(im).enhance(random.uniform(0.85, 1.15))
    w, h = im.size                            # leggero crop/zoom
    d = random.uniform(0, 0.08)
    im = im.crop((int(w*d), int(h*d), int(w*(1-d)), int(h*(1-d))))
    return im

def selftest(folder, n_aug=5):
    gallery = build_gallery(folder)
    names = list(gallery)
    print(f"Galleria: {len(names)} carte -> {names}")
    print("\nMatrice di similarita (carte diverse devono essere BEN separate):")
    print("      " + "  ".join(f"{n:>5}" for n in names))
    for a in names:
        row = "  ".join(f"{float(np.dot(gallery[a], gallery[b])):.2f}" for b in names)
        print(f"  {a:>3} {row}")
    print("\nTest match: per ogni carta, 5 versioni 'disturbate' -> ritrova sé stessa?")
    tot = ok = 0
    for name in names:
        base = Image.open(os.path.join(folder, name + os.path.splitext(glob.glob(os.path.join(folder,name+'*'))[0])[1]))
        for k in range(n_aug):
            q = embed(augment(base, k))
            top = match(q, gallery, topk=2)
            hit = top[0][1] == name
            tot += 1; ok += hit
            if not hit:
                print(f"  ❌ {name} aug{k}: 1°={top[0][1]}({top[0][0]:.2f}) 2°={top[1][1]}({top[1][0]:.2f})")
    print(f"\nAccuratezza match: {ok}/{tot} = {100*ok/tot:.0f}%")
    return ok == tot

if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--selftest", action="store_true")
    ap.add_argument("--gallery", default="recognition/gallery")
    ap.add_argument("--query")
    a = ap.parse_args()
    if a.selftest:
        sys.exit(0 if selftest(a.gallery) else 1)
    elif a.query:
        g = build_gallery(a.gallery)
        for sim, name in match(embed(Image.open(a.query)), g):
            print(f"  {name}: {sim:.3f}")
