"""
extract_palette.py

Extracts a primary + accent color from an uploaded logo (ignoring
transparent pixels and near-white/near-black/gray noise), derives a
darker "deep" variant of the primary, and generates monochrome +
reversed logo images for use in the brand manual's Logo Usage page.

No AI/ML dependency -- this is a straightforward alpha-aware color
bucketing pass, which works well because logos are almost always flat
vector-style art (a handful of true colors), not photographs.
"""

import base64
import colorsys
from collections import defaultdict
from io import BytesIO

import numpy as np
from PIL import Image


def _is_extreme(rgb):
    """True for near-white, near-black, or low-saturation gray midtones."""
    r, g, b = [c / 255 for c in rgb]
    mx, mn = max(r, g, b), min(r, g, b)
    l = (mx + mn) / 2
    s = 0 if mx == mn else (mx - mn) / (1 - abs(2 * l - 1))
    return l > 0.94 or l < 0.06 or (s < 0.08 and 0.1 < l < 0.9)


def _bucket_colors(pixels, bucket=12):
    buckets = defaultdict(list)
    for p in pixels:
        key = tuple((p // bucket) * bucket)
        buckets[key].append(p)
    clusters = []
    for plist in buckets.values():
        arr = np.array(plist)
        centroid = tuple(int(c) for c in arr.mean(axis=0))
        clusters.append((centroid, len(plist)))
    clusters.sort(key=lambda c: -c[1])
    return clusters


def _hue(rgb):
    r, g, b = [c / 255 for c in rgb]
    h, _, _ = colorsys.rgb_to_hsv(r, g, b)
    return h * 360


def _darken(rgb, factor=0.72):
    r, g, b = [c / 255 for c in rgb]
    h, l, s = colorsys.rgb_to_hls(r, g, b)
    l = max(0, l * factor)
    r2, g2, b2 = colorsys.hls_to_rgb(h, l, min(1, s * 1.05))
    return (int(r2 * 255), int(g2 * 255), int(b2 * 255))


def _to_hex(rgb):
    return "#{:02X}{:02X}{:02X}".format(*rgb)


def extract_palette(logo_path, max_colors=6):
    """Returns dict with primary/primary_deep/accent as hex + rgb-string pairs."""
    img = Image.open(logo_path).convert("RGBA")
    arr = np.array(img)
    alpha = arr[..., 3]

    mask = alpha > 200
    if mask.sum() < 50:
        mask = alpha > 10
    pixels = arr[mask][:, :3]

    if len(pixels) == 0:
        # fully transparent image edge case -- fall back to a neutral default
        pixels = np.array([[47, 75, 60]])

    filtered = np.array([p for p in pixels if not _is_extreme(p)])
    if len(filtered) < 0.05 * len(pixels):
        filtered = pixels  # brand genuinely is near-monochrome; don't over-filter

    clusters = _bucket_colors(filtered, bucket=12)[:max_colors]
    if not clusters:
        clusters = [((47, 75, 60), 1)]

    primary = clusters[0][0]

    accent = None
    for c, _ in clusters[1:]:
        if abs(_hue(c) - _hue(primary)) > 30 or abs(_hue(c) - _hue(primary)) < 330:
            if abs(((_hue(c) - _hue(primary)) + 180) % 360 - 180) > 30:
                accent = c
                break
    if accent is None:
        if len(clusters) > 1:
            accent = clusters[1][0]
        else:
            # near-monochrome logo: derive a complementary accent
            r, g, b = [c / 255 for c in primary]
            h, l, s = colorsys.rgb_to_hls(r, g, b)
            h2 = (h + 0.42) % 1.0
            r2, g2, b2 = colorsys.hls_to_rgb(h2, min(0.65, max(0.35, l)), max(0.45, s))
            accent = (int(r2 * 255), int(g2 * 255), int(b2 * 255))

    primary_deep = _darken(primary)

    def entry(rgb):
        return {"hex": _to_hex(rgb), "rgb": f"{rgb[0]},{rgb[1]},{rgb[2]}"}

    return {
        "primary": entry(primary),
        "primary_deep": entry(primary_deep),
        "accent": entry(accent),
    }


def recolor_logo(logo_path, hex_color):
    """Recolors every opaque pixel of the logo to a single flat color, keeping alpha."""
    img = Image.open(logo_path).convert("RGBA")
    arr = np.array(img)
    r = int(hex_color[1:3], 16)
    g = int(hex_color[3:5], 16)
    b = int(hex_color[5:7], 16)
    arr[..., 0] = r
    arr[..., 1] = g
    arr[..., 2] = b
    return Image.fromarray(arr, "RGBA")


def image_to_data_uri(img):
    buf = BytesIO()
    img.save(buf, format="PNG")
    b64 = base64.b64encode(buf.getvalue()).decode("ascii")
    return f"data:image/png;base64,{b64}"


def file_to_data_uri(path):
    with open(path, "rb") as f:
        b64 = base64.b64encode(f.read()).decode("ascii")
    ext = path.rsplit(".", 1)[-1].lower()
    mime = "jpeg" if ext == "jpg" else ext
    return f"data:image/{mime};base64,{b64}"
