#!/usr/bin/env python3
"""Auto-ritaglio carte FARAWAY da foto su sfondo uniforme (beige o feltro verde).

Rileva il rettangolo della carta come la regione che "stacca" dal colore di
sfondo (campionato dai bordi), in spazio Lab; prende il contorno piu grande,
ne ricava il rettangolo ruotato (minAreaRect) e raddrizza con warpPerspective.

Uso:
  python autocrop.py IN.jpg OUT.jpg            # una foto
  python autocrop.py --batch IN_DIR OUT_DIR    # tutte le *.jpg di una cartella
"""
import sys, os, glob, argparse
import numpy as np
import cv2

# Rapporto canonico = quello della carta di riferimento dell'utente (santuario.png:
# 2916x1908 => 1.528). Lo IMPONIAMO: tutte le carte hanno aspetto identico, verticale.
CANON_RATIO = 1.528
OUT_LONG = 1400
OUT_SHORT = int(round(OUT_LONG / CANON_RATIO))


def _localize(img):
    """bbox grezzo della carta (anche piccola nel frame): stacca dal fondo uniforme
    campionato ai bordi. Ritorna (x0,y0,x1,y1) full-res con padding, o None."""
    h, w = img.shape[:2]; sc = 900.0 / max(h, w)
    s = cv2.resize(img, None, fx=sc, fy=sc, interpolation=cv2.INTER_AREA)
    lab = cv2.cvtColor(s, cv2.COLOR_BGR2LAB).astype(np.float32)
    b = max(2, int(min(s.shape[:2]) * 0.05))
    bg = np.median(np.concatenate([lab[:b].reshape(-1, 3), lab[-b:].reshape(-1, 3),
                                   lab[:, :b].reshape(-1, 3), lab[:, -b:].reshape(-1, 3)]), 0)
    dn = np.linalg.norm(lab - bg, axis=2)
    dn = (dn / dn.max() * 255).astype(np.uint8)
    _, m = cv2.threshold(dn, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
    m = cv2.morphologyEx(m, cv2.MORPH_CLOSE, np.ones((11, 11), np.uint8), iterations=2)
    m = cv2.morphologyEx(m, cv2.MORPH_OPEN, np.ones((7, 7), np.uint8), iterations=1)
    cnts, _ = cv2.findContours(m, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
    if not cnts:
        return None
    c = max(cnts, key=cv2.contourArea)
    x, y, ww, hh = cv2.boundingRect(c)
    if ww * hh < 0.02 * s.shape[0] * s.shape[1]:
        return None
    pad = int(0.14 * max(ww, hh))
    x0, y0 = max(0, x - pad), max(0, y - pad)
    x1, y1 = min(s.shape[1], x + ww + pad), min(s.shape[0], y + hh + pad)
    return tuple(int(v / sc) for v in (x0, y0, x1, y1))


def find_card(img, debug=False):
    """Ritorna i 4 vertici (ordinati) della carta, o None.

    Strategia: 1) LOCALIZZA grossolanamente la carta (gestisce carte PICCOLE nel
    frame, con molto sfondo), 2) GrabCut sul ritaglio (cosi la carta riempie il box
    e GrabCut funziona anche con zone chiare/grigie), 3) 4 angoli via fit-delle-rette.
    """
    off = (0, 0)
    box = _localize(img)
    region = img
    if box is not None:
        x0, y0, x1, y1 = box
        if (x1 - x0) > 20 and (y1 - y0) > 20:
            region = img[y0:y1, x0:x1]; off = (x0, y0)

    h, w = region.shape[:2]
    scale = 700.0 / max(h, w)
    small = cv2.resize(region, None, fx=scale, fy=scale, interpolation=cv2.INTER_AREA)
    H, W = small.shape[:2]
    inset = 0.05
    rect = (int(W * inset), int(H * inset), int(W * (1 - 2 * inset)), int(H * (1 - 2 * inset)))
    mask = np.zeros((H, W), np.uint8)
    bgd = np.zeros((1, 65), np.float64); fgd = np.zeros((1, 65), np.float64)
    cv2.grabCut(small, mask, rect, bgd, fgd, 5, cv2.GC_INIT_WITH_RECT)
    m = np.where((mask == cv2.GC_FGD) | (mask == cv2.GC_PR_FGD), 255, 0).astype(np.uint8)
    m = cv2.morphologyEx(m, cv2.MORPH_OPEN, np.ones((5, 5), np.uint8))
    m = cv2.morphologyEx(m, cv2.MORPH_CLOSE, np.ones((9, 9), np.uint8), iterations=2)

    cnts, _ = cv2.findContours(m, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)
    if not cnts:
        return None
    c = max(cnts, key=cv2.contourArea)
    if cv2.contourArea(c) < 0.15 * H * W:
        return None
    quad = _quad_edgefit(c) / scale                      # -> coord. region full-res
    quad[:, 0] += off[0]; quad[:, 1] += off[1]           # -> coord. immagine full-res
    return _order_pts(quad)


def _line_intersect(l1, l2):
    (vx1, vy1, x1, y1) = l1; (vx2, vy2, x2, y2) = l2
    A = np.array([[vx1, -vx2], [vy1, -vy2]], float)
    b = np.array([x2 - x1, y2 - y1], float)
    t = np.linalg.solve(A, b)
    return np.array([x1 + vx1 * t[0], y1 + vy1 * t[0]], np.float32)


def _quad_edgefit(c):
    """4 angoli VERI = fit delle 4 RETTE dei bordi + intersezione.

    I lati dritti sono lunghi e affidabili; gli angoli arrotondati (rumorosi)
    vengono scartati. Molto piu preciso del vertex-picking di approxPolyDP, che
    a piccoli errori d'angolo "shear-ava" la carta (banda obliqua).
    """
    pts = c.reshape(-1, 2).astype(np.float32)
    box = cv2.boxPoints(cv2.minAreaRect(c))
    edges = [(box[i], box[(i + 1) % 4]) for i in range(4)]
    groups = [[] for _ in range(4)]
    for p in pts:                                         # assegna ogni punto al bordo piu vicino
        best, bd = 0, 1e18
        for gi, (a, b) in enumerate(edges):
            ab = b - a; t = np.clip(np.dot(p - a, ab) / (np.dot(ab, ab) + 1e-9), 0, 1)
            d = np.linalg.norm(p - (a + t * ab))
            if d < bd: bd, best = d, gi
        groups[best].append(p)
    lines = []
    for gi, (a, b) in enumerate(edges):
        g = np.array(groups[gi], np.float32)
        if len(g) < 10:
            lines.append(None); continue
        ab = b - a; L2 = np.dot(ab, ab) + 1e-9
        t = ((g - a) @ ab) / L2
        keep = g[(t > 0.12) & (t < 0.88)]                # via gli angoli arrotondati
        if len(keep) < 10: keep = g
        lines.append(tuple(cv2.fitLine(keep, cv2.DIST_L2, 0, 0.01, 0.01).ravel()))
    corners = []
    for i in range(4):
        l1, l2 = lines[i], lines[(i + 1) % 4]
        corners.append(box[(i + 1) % 4] if (l1 is None or l2 is None) else _line_intersect(l1, l2))
    return np.array(corners, np.float32)


def _order_pts(pts):
    pts = np.array(pts, dtype=np.float32)
    s = pts.sum(axis=1); d = np.diff(pts, axis=1).ravel()
    return np.array([pts[np.argmin(s)],   # top-left
                     pts[np.argmin(d)],   # top-right
                     pts[np.argmax(s)],   # bottom-right
                     pts[np.argmax(d)]],  # bottom-left
                    dtype=np.float32)


def _band_angle(card):
    """Angolo della BANDA = la riga a massima energia di bordo orizzontale (full-width,
    piu forte di icone/orizzonti). Ritorna None se poco affidabile."""
    g = cv2.cvtColor(card, cv2.COLOR_BGR2GRAY).astype(np.float32)
    H, W = g.shape
    E = cv2.GaussianBlur(np.abs(cv2.Sobel(g, cv2.CV_32F, 0, 1, ksize=3)), (0, 0), 1.5)
    y0, y1 = int(H * 0.40), int(H * 0.82)
    rowE = cv2.GaussianBlur(E[y0:y1].sum(1).reshape(-1, 1), (0, 0), 3).ravel()
    yc = y0 + int(np.argmax(rowE))
    win = int(H * 0.06)
    xs, ys = [], []
    for x in range(int(W * 0.06), int(W * 0.94), 2):
        seg = E[max(0, yc - win):yc + win, x]
        j = int(np.argmax(seg))
        if seg[j] > seg.mean() * 1.5:
            xs.append(x); ys.append(max(0, yc - win) + j)
    if len(xs) < 25:
        return None
    vx, vy, _, _ = cv2.fitLine(np.column_stack([np.array(xs, np.float32), np.array(ys, np.float32)]),
                               cv2.DIST_HUBER, 0, 0.01, 0.01).ravel()
    a = np.degrees(np.arctan2(vy, vx))
    return a - 180 if a > 90 else (a + 180 if a < -90 else a)


def _level_band(card):
    """Ruota la carta per rendere la BANDA orizzontale. Guardia 'non peggiorare':
    se dopo la rotazione il residuo non e piu piccolo, annulla."""
    a = _band_angle(card)
    if a is None or abs(a) < 0.2 or abs(a) > 8:
        return card
    h, w = card.shape[:2]
    M = cv2.getRotationMatrix2D((w / 2, h / 2), a, 1.0)
    rot = cv2.warpAffine(card, M, (w, h), flags=cv2.INTER_CUBIC, borderMode=cv2.BORDER_REPLICATE)
    a2 = _band_angle(rot)
    return rot if (a2 is not None and abs(a2) < abs(a)) else card


def _panel_centroid_y(bgr):
    """Y (frazione 0..1) del baricentro della piu grande regione di colore PIATTO
    connessa = il pannello bioma. Robusto: il pannello e UN blob grande (l'arte e
    texturata), a differenza dei pixel piatti sparsi."""
    g = cv2.cvtColor(bgr, cv2.COLOR_BGR2GRAY).astype(np.float32)
    h = g.shape[0]
    grad = np.abs(cv2.Sobel(g, cv2.CV_32F, 1, 0, 3)) + np.abs(cv2.Sobel(g, cv2.CV_32F, 0, 1, 3))
    grad = cv2.GaussianBlur(grad, (0, 0), 2)
    flat = cv2.morphologyEx((grad < 14).astype(np.uint8), cv2.MORPH_OPEN, np.ones((9, 9), np.uint8))
    n, lab, st, ce = cv2.connectedComponentsWithStats(flat, 8)
    if n <= 1:
        return 0.5
    k = max(range(1, n), key=lambda i: st[i, cv2.CC_STAT_AREA])
    return ce[k][1] / h


def _auto_orient(card):
    """Verso canonico: PORTRAIT con illustrazione in ALTO e pannello bioma in BASSO.
    Decide il 180 dal baricentro del pannello (deve stare nella meta bassa)."""
    if card.shape[1] > card.shape[0]:          # se landscape -> portrait
        card = cv2.rotate(card, cv2.ROTATE_90_CLOCKWISE)
    if _panel_centroid_y(card) < 0.5:          # pannello in alto -> ribalta
        card = cv2.rotate(card, cv2.ROTATE_180)
    return card


def warp(img, quad, margin=0):
    """Raddrizza, IMPONE il rapporto canonico e l'orientamento (PORTRAIT, arte in alto)."""
    (tl, tr, br, bl) = quad
    wlen = (np.linalg.norm(tr - tl) + np.linalg.norm(br - bl)) / 2
    hlen = (np.linalg.norm(bl - tl) + np.linalg.norm(br - tr)) / 2
    if hlen >= wlen:                       # lato lungo gia verticale
        W, H = OUT_SHORT, OUT_LONG
    else:                                  # lato lungo orizzontale
        W, H = OUT_LONG, OUT_SHORT
    dst = np.array([[0, 0], [W - 1, 0], [W - 1, H - 1], [0, H - 1]], dtype=np.float32)
    out = cv2.warpPerspective(img, cv2.getPerspectiveTransform(quad, dst), (W, H))
    return _level_band(_auto_orient(out))


def autocrop(in_path, out_path):
    img = cv2.imread(in_path)
    if img is None:
        print(f"  !! impossibile leggere {in_path}"); return False
    quad = find_card(img)
    if quad is None:
        print(f"  !! carta non trovata in {os.path.basename(in_path)}"); return False
    card = warp(img, quad, margin=4)
    cv2.imwrite(out_path, card, [cv2.IMWRITE_JPEG_QUALITY, 95])
    print(f"  ok {os.path.basename(in_path)} -> {os.path.basename(out_path)} "
          f"({card.shape[1]}x{card.shape[0]})")
    return True


if __name__ == "__main__":
    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")) +
                        glob.glob(os.path.join(ind, "*.jpeg")) +
                        glob.glob(os.path.join(ind, "*.png"))):
            autocrop(p, os.path.join(outd, os.path.splitext(os.path.basename(p))[0] + ".jpg"))
    elif len(a.paths) == 2:
        autocrop(a.paths[0], a.paths[1])
    else:
        ap.print_help()
