#!/usr/bin/env python3
"""Rifinitura carta per l'APP: dal ritaglio raddrizzato (autocrop, gia PORTRAIT e a
rapporto canonico) -> PNG RGBA con sfondo TRASPARENTE e angoli arrotondati IDENTICI
allo schema dell'utente.

La maschera angoli = il canale alpha della carta di riferimento `recognition/card_mask.png`
(estratto da santuario.png), ridimensionato all'uscita. Cosi i bordi sono esattamente
quelli che l'utente vuole.
"""
import os, glob, cv2, numpy as np

OUT_H = 1400          # altezza uniforme (portrait); larghezza segue il rapporto della carta
TRIM_FRAC = 0.010     # rosicchia un filo verso l'interno (toglie eventuale frangia di tavolo)
MASK_PATH = os.path.join(os.path.dirname(__file__), "card_mask.png")
TARGET_L = 150        # luminosità mediana target (uniforma tutte le carte)

_mask_cache = {}


def normalize_light(bgr):
    """Uniforma luce+colore SENZA bruciare le alte luci: (1) bilancia il bianco sui
    pixel NEUTRI (robusto ai riflessi), (2) porta la mediana di luminosità a TARGET_L
    con una CURVA GAMMA (non un moltiplicatore lineare): solleva mezzitoni/ombre ma il
    bianco resta <=255, quindi il celeste del cielo non diventa bianco."""
    f = bgr.astype(np.float32)
    S = cv2.cvtColor(bgr, cv2.COLOR_BGR2HSV)[..., 1]
    neu = S < 45
    if neu.sum() > bgr.shape[0] * bgr.shape[1] * 0.01:
        ref = f[neu].reshape(-1, 3).mean(0)
        gains = np.clip(ref.mean() / np.clip(ref, 1, None), 0.6, 1.6)   # WB dolce
        f = np.clip(f * gains, 0, 255)
    # luminosità: gamma che porta la mediana a TARGET_L (preserva le alte luci)
    med = np.median(cv2.cvtColor(f.astype(np.uint8), cv2.COLOR_BGR2GRAY))
    med = min(max(med, 5), 250)
    gamma = np.log(TARGET_L / 255.0) / np.log(med / 255.0)
    gamma = float(np.clip(gamma, 0.45, 1.8))
    lut = (((np.arange(256) / 255.0) ** gamma) * 255).astype(np.uint8)
    return lut[f.astype(np.uint8)]


def _mask(h, w):
    key = (h, w)
    if key not in _mask_cache:
        m = cv2.imread(MASK_PATH, cv2.IMREAD_GRAYSCALE)
        _mask_cache[key] = cv2.resize(m, (w, h), interpolation=cv2.INTER_AREA)
    return _mask_cache[key]


def finish(in_path, out_path):
    im = cv2.imread(in_path)
    im = normalize_light(im)
    h, w = im.shape[:2]
    tx, ty = int(w * TRIM_FRAC), int(h * TRIM_FRAC)
    im = im[ty:h - ty, tx:w - tx]
    # normalizza altezza mantenendo il rapporto
    h, w = im.shape[:2]
    s = OUT_H / h
    im = cv2.resize(im, (int(round(w * s)), OUT_H), interpolation=cv2.INTER_AREA)
    h, w = im.shape[:2]
    rgba = cv2.cvtColor(im, cv2.COLOR_BGR2BGRA)
    rgba[:, :, 3] = _mask(h, w)
    cv2.imwrite(out_path, rgba)
    return out_path


if __name__ == "__main__":
    import argparse
    ap = argparse.ArgumentParser()
    ap.add_argument("--batch", nargs=2, metavar=("IN_DIR", "OUT_DIR"))
    ap.add_argument("paths", nargs="*")
    a = ap.parse_args()
    if a.batch:
        ind, outd = a.batch; os.makedirs(outd, exist_ok=True)
        for p in sorted(glob.glob(os.path.join(ind, "*.jpg"))):
            o = os.path.join(outd, os.path.splitext(os.path.basename(p))[0] + ".png")
            finish(p, o); print("ok", os.path.basename(o))
    elif len(a.paths) == 2:
        finish(a.paths[0], a.paths[1]); print("ok", a.paths[1])
    else:
        ap.print_help()
