#!/usr/bin/env python3
"""LANDR PRO MASTER — live DJ insert. Pedalboard + meters from the same buffer. No VST."""

from __future__ import annotations

import json
import math
import sys
import threading
import time
from pathlib import Path

import numpy as np
import pygame
import sounddevice as sd
from pedalboard import (
    Compressor,
    Distortion,
    Gain,
    HighpassFilter,
    HighShelfFilter,
    Limiter,
    LowShelfFilter,
    PeakFilter,
    Pedalboard,
)

ROOT = Path(__file__).resolve().parent
STATE_PATH = ROOT / "state.json"

WANT_IN = ("BlackHole 2ch", "BlackHole")
WANT_OUT = ("External Headphones", "MacBook Air Speakers")
FORBID_OUT = (
    "BlackHole", "Aggregate", "Multi-Output", "Multi-Out", "DJ Viz",
    "01 DJ PA", "PA+ASUS",
)
STYLES = ("WARM", "BALANCED", "OPEN")
BARS = 64
FFT_N = 1024
DETECT_DB = -48.0
# DJ-set bands — numbered, not color-only
DJ_BANDS = (
    ("01 SUB", 20.0, 80.0),
    ("02 BASS", 80.0, 200.0),
    ("03 BODY", 200.0, 500.0),
    ("04 MID", 500.0, 1800.0),
    ("05 VOC", 1800.0, 3500.0),
    ("06 HI", 3500.0, 8000.0),
    ("07 AIR", 8000.0, 16000.0),
)

C_BG = (7, 5, 20)
C_PANEL = (20, 14, 48)
C_WELL = (10, 7, 28)
C_TEAL = (129, 140, 248)
C_TEAL_DIM = (67, 56, 140)
C_ORANGE = (251, 146, 60)
C_SKY = (165, 180, 255)
C_GREEN = (94, 234, 212)
C_YELLOW = (196, 181, 253)
C_RED = (244, 114, 182)
C_MUTED = (148, 140, 180)
C_WHITE = (236, 233, 255)
C_KNOB = (32, 24, 64)
C_VIOLET = (167, 139, 250)
C_DEEP = (49, 32, 120)

DEFAULTS = {
    "in_gain": 0.0,
    "style": 1,
    "eq_low": 0.0,
    "eq_mid": 0.0,
    "eq_high": 0.0,
    "presence": -0.8,
    "deess": 0.0,
    "width": 0.50,
    "comp": 0.50,
    "character": 0.28,
    "sat": 0.0,
    "loudness": 0.54,
    "bypass": False,
    "gain_match": False,
}


def db(x: float) -> float:
    return 20.0 * math.log10(max(x, 1e-12))


def sd_pick(kind: str, wants: tuple[str, ...], forbid: tuple[str, ...] = ()) -> tuple[int, str]:
    fallback = None
    for i, d in enumerate(sd.query_devices()):
        chans = d["max_input_channels"] if kind == "in" else d["max_output_channels"]
        if chans < 1:
            continue
        name = d["name"]
        if any(f.lower() in name.lower() for f in forbid):
            continue
        if fallback is None:
            fallback = (i, name)
        for want in wants:
            if want.lower() in name.lower():
                return i, name
    if fallback:
        return fallback
    sys.exit("No audio devices.")


def resample_lin(x: np.ndarray, src: int, dst: int) -> np.ndarray:
    if src == dst or x.shape[0] < 2:
        return x
    n = max(1, int(round(x.shape[0] * dst / src)))
    t = np.linspace(0.0, 1.0, x.shape[0], endpoint=False)
    tn = np.linspace(0.0, 1.0, n, endpoint=False)
    y = np.empty((n, x.shape[1]), dtype=np.float32)
    for c in range(x.shape[1]):
        y[:, c] = np.interp(tn, t, x[:, c])
    return y


class Chain:
    def __init__(self) -> None:
        self.in_gain = DEFAULTS["in_gain"]
        self.style = DEFAULTS["style"]
        self.eq_low = DEFAULTS["eq_low"]
        self.eq_mid = DEFAULTS["eq_mid"]
        self.eq_high = DEFAULTS["eq_high"]
        self.presence = DEFAULTS["presence"]
        self.deess = DEFAULTS["deess"]
        self.width = DEFAULTS["width"]
        self.comp = DEFAULTS["comp"]
        self.character = DEFAULTS["character"]
        self.sat = DEFAULTS["sat"]
        self.loudness = DEFAULTS["loudness"]
        self.bypass = False
        self.gain_match = False
        self.match_db = 0.0
        self.g_in = Gain(gain_db=0.0)
        self.hpf = HighpassFilter(cutoff_frequency_hz=22.0)
        self.comp_p = Compressor(threshold_db=-16.0, ratio=2.8, attack_ms=8.0, release_ms=90.0)
        self.low = LowShelfFilter(cutoff_frequency_hz=90.0, gain_db=0.6)
        self.mid = PeakFilter(cutoff_frequency_hz=400.0, gain_db=0.0, q=0.55)
        self.pres = PeakFilter(cutoff_frequency_hz=2200.0, gain_db=-0.8, q=0.80)
        self.deess_p = PeakFilter(cutoff_frequency_hz=7200.0, gain_db=0.0, q=1.10)
        self.high = HighShelfFilter(cutoff_frequency_hz=9000.0, gain_db=0.4)
        self.air = HighShelfFilter(cutoff_frequency_hz=12000.0, gain_db=0.0)
        self.dist = Distortion(drive_db=0.0)
        self.g_loud = Gain(gain_db=2.5)
        self.g_match = Gain(gain_db=0.0)
        self.lim = Limiter(threshold_db=-1.8, release_ms=80.0)
        self.parts = [
            self.g_in, self.hpf, self.comp_p, self.low, self.mid, self.pres,
            self.deess_p, self.high, self.air, self.dist, self.g_loud, self.g_match, self.lim,
        ]
        self.board = Pedalboard(self.parts)
        self.apply()

    def apply_dict(self, d: dict) -> None:
        for k in DEFAULTS:
            if k in d and k not in ("bypass", "gain_match", "style"):
                setattr(self, k, float(d[k]))
        if "style" in d:
            self.style = int(d["style"]) % 3
        if "bypass" in d:
            self.bypass = bool(d["bypass"])
        if "gain_match" in d:
            self.gain_match = bool(d["gain_match"])
        self.apply()

    def dump(self) -> dict:
        return {
            "in_gain": self.in_gain,
            "style": self.style,
            "eq_low": self.eq_low,
            "eq_mid": self.eq_mid,
            "eq_high": self.eq_high,
            "presence": self.presence,
            "deess": self.deess,
            "width": self.width,
            "comp": self.comp,
            "character": self.character,
            "sat": self.sat,
            "loudness": self.loudness,
            "bypass": self.bypass,
            "gain_match": self.gain_match,
        }

    def apply(self) -> None:
        c = float(np.clip(self.comp, 0.0, 1.0))
        ch = float(np.clip(self.character, 0.0, 1.0))
        w = float(np.clip(self.width, 0.0, 1.0))
        s = float(np.clip(self.sat, 0.0, 1.0))
        loud = -4.0 + 8.0 * float(np.clip(self.loudness, 0.0, 1.0))
        attack = 3.0 + 18.0 * ch
        release = 40.0 + 179.0 * ch
        low = self.eq_low
        mid = self.eq_mid
        high = self.eq_high
        if self.style == 0:
            low += 1.4
            high -= 0.2
            attack += 6.0
        elif self.style == 2:
            low += 0.2
            high += 1.2
            attack -= 2.0
        else:
            low += 0.6
            high += 0.4
        self.g_in.gain_db = self.in_gain
        self.comp_p.threshold_db = -8.0 - 16.0 * c
        self.comp_p.ratio = 1.6 + 2.4 * c
        self.comp_p.attack_ms = max(2.0, attack)
        self.comp_p.release_ms = release
        self.low.gain_db = low
        self.mid.gain_db = mid
        self.pres.gain_db = self.presence
        self.deess_p.gain_db = -8.0 * float(np.clip(self.deess, 0.0, 1.0))
        self.high.gain_db = high
        self.air.gain_db = (w - 0.5) * 2.4
        self.dist.drive_db = 0.0 if s < 0.02 else 10.0 * s
        self.g_loud.gain_db = loud
        self.g_match.gain_db = float(np.clip(self.match_db, -4.0, 4.0)) if self.gain_match else 0.0

    def process(self, block: np.ndarray, sr: float) -> np.ndarray:
        if self.bypass:
            return np.clip(block, -0.89, 0.89)
        x = np.ascontiguousarray(block.T)
        if float(np.clip(self.sat, 0.0, 1.0)) < 0.02:
            self.dist.drive_db = 0.0
        y = self.board(x, float(sr))
        if y.ndim == 1:
            y = np.stack([y, y], axis=0)
        y = np.clip(np.nan_to_num(y, copy=False), -0.89, 0.89)
        return np.ascontiguousarray(y.T)


class Live:
    def __init__(self, ch: Chain, in_dev: int, out_dev: int, sr: int) -> None:
        self.ch = ch
        self.sr = 48000
        self.out_sr = 48000
        self.lock = threading.Lock()
        self.buf = np.zeros(FFT_N, np.float32)
        self.err = ""
        cap = 48000
        self.ring = np.zeros((cap, 2), np.float32)
        self.r = self.w = self.n = 0
        kw = dict(blocksize=512, channels=2, dtype="float32", latency="low")
        self.sin = sd.InputStream(samplerate=48000, device=in_dev, callback=self._in, **kw)
        try:
            self.sout = sd.OutputStream(samplerate=48000, device=out_dev, callback=self._out, **kw)
        except Exception:
            self.out_sr = int(sd.query_devices(out_dev).get("default_samplerate") or 44100)
            self.sout = sd.OutputStream(samplerate=self.out_sr, device=out_dev, callback=self._out, **kw)
        self.sin.start()
        self.sout.start()

    def _push(self, y: np.ndarray) -> None:
        n = int(y.shape[0])
        cap = self.ring.shape[0]
        if n <= 0:
            return
        if n >= cap:
            self.ring[:] = y[-cap:]
            self.r = 0
            self.w = 0
            self.n = cap
            return
        free = cap - self.n
        if n > free:
            drop = n - free
            self.r = (self.r + drop) % cap
            self.n -= drop
        end = self.w + n
        if end <= cap:
            self.ring[self.w:end] = y
        else:
            k = cap - self.w
            self.ring[self.w:] = y[:k]
            self.ring[: n - k] = y[k:]
        self.w = (self.w + n) % cap
        self.n += n

    def _pop(self, frames: int) -> np.ndarray:
        out = np.zeros((frames, 2), np.float32)
        n = min(frames, self.n)
        if n <= 0:
            return out
        cap = self.ring.shape[0]
        end = self.r + n
        if end <= cap:
            out[:n] = self.ring[self.r:end]
        else:
            k = cap - self.r
            out[:k] = self.ring[self.r:]
            out[k:n] = self.ring[: n - k]
        self.r = (self.r + n) % cap
        self.n -= n
        return out

    def _in(self, indata, frames, time_info, status) -> None:
        xin = np.nan_to_num(np.asarray(indata, dtype=np.float32), copy=False)
        if xin.ndim == 1:
            xin = np.column_stack([xin, xin])
        xin = np.clip(xin[:, :2], -1.0, 1.0)
        try:
            y = self.ch.process(xin, float(self.sr))
        except Exception:
            y = xin
        y = np.clip(np.nan_to_num(np.asarray(y, dtype=np.float32), copy=False), -1.0, 1.0)
        if y.ndim == 1:
            y = np.column_stack([y, y])
        if self.out_sr != self.sr:
            y = resample_lin(y, self.sr, self.out_sr)
        mono = xin.mean(axis=1)
        k = min(len(mono), FFT_N)
        with self.lock:
            self._push(y[:, :2])
            self.buf[:-k] = self.buf[k:]
            self.buf[-k:] = mono[-k:]

    def _out(self, outdata, frames, time_info, status) -> None:
        outdata.fill(0)
        with self.lock:
            got = self._pop(frames)
        n = min(frames, got.shape[0], outdata.shape[0])
        outdata[:n, :2] = got[:n, :2]

    def snap(self) -> np.ndarray:
        with self.lock:
            return self.buf.copy()

    def close(self) -> None:
        for s in (getattr(self, "sin", None), getattr(self, "sout", None)):
            if s is None:
                continue
            try:
                s.stop()
                s.close()
            except Exception:
                pass


class Knob:
    def __init__(self, key: str, num: str, name: str, lo: float, hi: float, getter, setter, unit: str):
        self.key, self.num, self.name = key, num, name
        self.lo, self.hi, self.unit = lo, hi, unit
        self.getter, self.setter = getter, setter
        self.rect = pygame.Rect(0, 0, 96, 118)

    def place(self, x: int, y: int) -> None:
        self.rect.topleft = (x, y)

    def hit(self, pos) -> bool:
        return self.rect.collidepoint(pos)

    def nudge(self, dn: float) -> None:
        span = self.hi - self.lo
        self.setter(float(np.clip(self.getter() + dn * span, self.lo, self.hi)))

    def draw(self, surf, font, small) -> None:
        x, y, w, h = self.rect
        pygame.draw.rect(surf, C_PANEL, self.rect, border_radius=10)
        cx, cy, r = x + w // 2, y + 44, 28
        pygame.draw.circle(surf, C_KNOB, (cx, cy), r)
        pygame.draw.circle(surf, C_WELL, (cx, cy), r - 5)
        n = (self.getter() - self.lo) / (self.hi - self.lo)
        ang = math.radians(220 - 260 * n)
        pygame.draw.line(
            surf, C_TEAL, (cx, cy),
            (cx + int(math.cos(ang) * (r - 8)), cy - int(math.sin(ang) * (r - 8))), 3,
        )
        pygame.draw.circle(surf, C_WHITE, (cx, cy), 3)
        surf.blit(small.render(f"{self.num} {self.name}", True, C_MUTED), (x + 6, y + 78))
        val = self.getter()
        if self.unit == "%":
            txt = f"{val * 100:.0f}%"
        elif self.unit == "dB":
            txt = f"{val:+.1f} dB"
        else:
            txt = f"{val:.2f}"
        surf.blit(font.render(txt, True, C_WHITE), (x + 8, y + 94))


def draw_vumeter(surf, font, small, x, y, w, h, value_db, hold_db, title, color):
    pygame.draw.rect(surf, C_PANEL, (x, y, w, h), border_radius=10)
    surf.blit(small.render(title, True, C_MUTED), (x + 10, y + 8))
    inner = pygame.Rect(x + 18, y + 36, w - 36, h - 70)
    pygame.draw.rect(surf, C_WELL, inner, border_radius=4)
    n = float(np.clip((value_db + 60.0) / 60.0, 0.0, 1.0))
    hn = float(np.clip((hold_db + 60.0) / 60.0, 0.0, 1.0))
    fh = int(inner.h * n)
    col = C_RED if value_db > -1.2 else color
    pygame.draw.rect(surf, col, (inner.x + 4, inner.bottom - fh, inner.w - 8, fh), border_radius=3)
    pygame.draw.rect(surf, C_YELLOW, (inner.x + 4, inner.bottom - int(inner.h * hn), inner.w - 8, 3))
    surf.blit(font.render(f"{value_db:5.1f}", True, C_WHITE), (x + 10, y + h - 28))


class TypePool:
    def __init__(self) -> None:
        self.cache: dict[int, pygame.font.Font] = {}

    def get(self, size: int) -> pygame.font.Font:
        size = int(np.clip(size, 10, 92))
        if size not in self.cache:
            self.cache[size] = pygame.font.Font(None, size)
        return self.cache[size]


def ink(t: float, energy: float) -> tuple[int, int, int]:
    t = float(np.clip(t, 0.0, 1.0))
    e = float(np.clip(energy, 0.0, 1.0))
    r = int(70 + 90 * t + 60 * e)
    g = int(40 + 50 * (1 - t) + 40 * e)
    b = int(160 + 80 * (1 - t) + 15 * e)
    return (min(255, r), min(255, g), min(255, b))


def liquid_word(
    surf, pool: TypePool, text: str, x: int, y: int, base: int,
    energy: float, wave: np.ndarray, t: float, tracking: int = 1,
    measuring: bool = False,
) -> None:
    if not measuring:
        glyph = pool.get(base).render(text, True, C_MUTED)
        surf.blit(glyph, (x, y))
        return
    cx = x
    n = max(len(text) - 1, 1)
    for i, ch in enumerate(text):
        if ch == " ":
            cx += int(base * 0.32)
            continue
        lift = 0.0
        if wave.size:
            wi = float(wave[int(i / n * (wave.size - 1))])
            lift = wi * (16 + 40 * energy)
        sz = int(base * (1.0 + 0.22 * energy))
        col = ink(i / n, energy)
        glyph = pool.get(sz).render(ch, True, col)
        glow = pool.get(sz).render(ch, True, C_DEEP)
        gy = y + int(lift)
        surf.blit(glow, (cx + 2, gy + 2))
        surf.blit(glyph, (cx, gy))
        cx += glyph.get_width() + tracking


def make_bg(w: int, h: int) -> pygame.Surface:
    s = pygame.Surface((w, h))
    for y in range(h):
        u = y / max(h - 1, 1)
        s.fill((int(6 + 10 * u), int(3 + 6 * (1 - u)), int(18 + 22 * (1 - u) + 8 * u)), (0, y, w, 1))
    return s


def band_levels(mono: np.ndarray, sr: float) -> np.ndarray:
    out = np.zeros(len(DJ_BANDS), np.float32)
    if len(mono) < 32:
        return out
    n = min(len(mono), FFT_N)
    mag = np.abs(np.fft.rfft(mono[-n:] * np.hanning(n))) + 1e-12
    freqs = np.fft.rfftfreq(n, 1.0 / sr)
    for i, (_name, lo, hi) in enumerate(DJ_BANDS):
        m = (freqs >= lo) & (freqs < hi)
        if np.any(m):
            out[i] = float(np.clip((db(float(np.mean(mag[m]))) + 80.0) / 80.0, 0.0, 1.0))
    return out


def draw_liquid(
    surf, rect: pygame.Rect, spec: np.ndarray, wave: np.ndarray,
    energy: float, live: bool, bands: np.ndarray, dominant: int, flash: str,
    pool: TypePool, small, font,
) -> None:
    well = pygame.Surface(rect.size, pygame.SRCALPHA)
    well.fill((8, 4, 24, 235) if live else (12, 10, 16, 240))
    if spec.size >= 4:
        k = 7
        pad = np.pad(spec, k // 2, mode="edge")
        ys = np.convolve(pad, np.ones(k) / k, mode="valid")
        xs = np.linspace(18, rect.w - 18, len(ys))
        base = rect.h - 86
        mh = rect.h - 130
        top = [(int(xs[0]), base)]
        for i, v in enumerate(ys):
            top.append((int(xs[i]), int(base - mh * float(v))))
        top.append((int(xs[-1]), base))
        if len(top) > 4:
            body = (48, 30, 140, 180) if live else (40, 40, 48, 90)
            ridge = (210, 200, 255, 230) if live else (110, 110, 120, 100)
            pygame.draw.polygon(well, body, top)
            if len(top) > 3:
                pygame.draw.lines(well, ridge, False, top[1:-1], 2 if live else 1)
    if live and wave.size > 8:
        mid = 110
        step = max(1, wave.size // 200)
        wpts = []
        for j in range(0, wave.size, step):
            wx = 18 + (rect.w - 36) * j / (wave.size - 1)
            wy = mid + float(wave[j]) * (20 + 40 * energy)
            wpts.append((int(wx), int(np.clip(wy, 8, rect.h - 96))))
        if len(wpts) > 1:
            pygame.draw.lines(well, (170, 150, 255, 210), False, wpts, 2)
    bw = (rect.w - 36) // len(DJ_BANDS)
    for i, (name, _lo, _hi) in enumerate(DJ_BANDS):
        x = 18 + i * bw
        h = int(54 * float(bands[i] if i < len(bands) else 0))
        on = live and i == dominant
        col = C_TEAL if on else (C_TEAL_DIM if live else (50, 50, 58))
        pygame.draw.rect(well, (16, 12, 32), (x, rect.h - 72, bw - 8, 58), border_radius=6)
        if h > 2:
            pygame.draw.rect(well, col, (x + 4, rect.h - 18 - h, bw - 16, h), border_radius=3)
        lab = pool.get(15).render(name, True, C_WHITE if on else C_MUTED)
        well.blit(lab, (x + 4, rect.h - 70))
    surf.blit(well, rect.topleft)
    pygame.draw.rect(surf, C_GREEN if live else C_ORANGE, rect, 3 if live else 2, border_radius=16)
    if live:
        now = DJ_BANDS[dominant][0] if 0 <= dominant < len(DJ_BANDS) else "—"
        surf.blit(font.render(f"22 LIVE FREQ   measuring   NOW {now}", True, C_GREEN), (rect.x + 20, rect.y + 10))
        if flash:
            surf.blit(font.render(f"23 MOVE  {flash}", True, C_YELLOW), (rect.x + 520, rect.y + 10))
    else:
        surf.blit(font.render("22 STATIC   not measuring frequency", True, C_ORANGE), (rect.x + 20, rect.y + 10))
        surf.blit(small.render("last shape frozen  ·  djay must hit BlackHole to measure", True, C_MUTED), (rect.x + 20, rect.y + 34))


def spectrum(mono: np.ndarray, sr: float) -> np.ndarray:
    if len(mono) < 32:
        return np.zeros(BARS, np.float32)
    n = min(len(mono), FFT_N)
    x = mono[-n:] * np.hanning(n)
    mag = np.abs(np.fft.rfft(x)) + 1e-12
    freqs = np.fft.rfftfreq(n, 1.0 / sr)
    edges = np.geomspace(30.0, 16000.0, BARS + 1)
    bars = np.zeros(BARS, np.float32)
    for i in range(BARS):
        m = (freqs >= edges[i]) & (freqs < edges[i + 1])
        if np.any(m):
            bars[i] = (db(float(np.mean(mag[m]))) + 80.0) / 80.0
    return np.clip(bars, 0.0, 1.0)


def main() -> None:
    pygame.init()
    pygame.display.set_caption("01 LANDR PRO MASTER")
    screen = pygame.display.set_mode((1400, 820))
    clock = pygame.time.Clock()
    pool = TypePool()
    font = pool.get(22)
    small = pool.get(16)
    huge = pool.get(48)
    big = pool.get(32)
    bg = make_bg(1400, 820)
    screen.blit(bg, (0, 0))
    liquid_word(screen, pool, "LANDR PRO MASTER", 40, 36, 48, 0.2, np.zeros(8), 0.0)
    screen.blit(font.render("starting…", True, C_VIOLET), (40, 100))
    pygame.display.flip()
    pygame.event.pump()

    in_dev, inn = sd_pick("in", WANT_IN)
    out_dev, out = sd_pick("out", WANT_OUT, forbid=FORBID_OUT)
    ch = Chain()
    if STATE_PATH.exists():
        try:
            saved = json.loads(STATE_PATH.read_text())
            ch.apply_dict(saved)
        except Exception:
            pass
    sr = int(sd.query_devices(in_dev).get("default_samplerate") or 48000)
    live = None
    err = ""
    try:
        live = Live(ch, in_dev, out_dev, sr)
    except Exception as e:
        err = str(e)

    def persist() -> None:
        STATE_PATH.write_text(json.dumps(ch.dump(), indent=2))

    def set_bypass(on: bool) -> None:
        ch.bypass = on

    knobs = [
        Knob("in_gain", "01", "IN GAIN", -12, 12, lambda: ch.in_gain, lambda v: (setattr(ch, "in_gain", v), ch.apply()), "dB"),
        Knob("eq_low", "03", "EQ LOW", -6, 6, lambda: ch.eq_low, lambda v: (setattr(ch, "eq_low", v), ch.apply()), "dB"),
        Knob("eq_mid", "04", "EQ MID", -6, 6, lambda: ch.eq_mid, lambda v: (setattr(ch, "eq_mid", v), ch.apply()), "dB"),
        Knob("eq_high", "05", "EQ HIGH", -6, 6, lambda: ch.eq_high, lambda v: (setattr(ch, "eq_high", v), ch.apply()), "dB"),
        Knob("presence", "06", "PRESENCE", -6, 6, lambda: ch.presence, lambda v: (setattr(ch, "presence", v), ch.apply()), "dB"),
        Knob("deess", "07", "DE-ESS", 0, 1, lambda: ch.deess, lambda v: (setattr(ch, "deess", v), ch.apply()), "%"),
        Knob("width", "08", "WIDTH", 0, 1, lambda: ch.width, lambda v: (setattr(ch, "width", v), ch.apply()), "%"),
        Knob("comp", "09", "COMP", 0, 1, lambda: ch.comp, lambda v: (setattr(ch, "comp", v), ch.apply()), "%"),
        Knob("character", "10", "CHAR", 0, 1, lambda: ch.character, lambda v: (setattr(ch, "character", v), ch.apply()), "%"),
        Knob("sat", "11", "SAT", 0, 1, lambda: ch.sat, lambda v: (setattr(ch, "sat", v), ch.apply()), "%"),
        Knob("loudness", "12", "LOUDNESS", 0, 1, lambda: ch.loudness, lambda v: (setattr(ch, "loudness", v), ch.apply()), "%"),
    ]
    for i, k in enumerate(knobs):
        k.place(24 + i * 108, 668)
    style_rects = [pygame.Rect(24 + i * 120, 612, 110, 40) for i in range(3)]
    btn_match = pygame.Rect(980, 28, 130, 40)
    btn_bypass = pygame.Rect(1120, 28, 110, 40)
    spec = np.zeros(BARS, np.float32)
    hold = np.zeros(BARS, np.float32)
    wave = np.zeros(512, np.float32)
    bands = np.zeros(len(DJ_BANDS), np.float32)
    prev_bands = np.zeros(len(DJ_BANDS), np.float32)
    frozen_spec = spec.copy()
    frozen_bands = bands.copy()
    dominant = 0
    flash = ""
    flash_until = 0.0
    in_rms = out_rms = 1e-6
    in_peak = out_peak = 1e-6
    lufs = tp_hold = -70.0
    detected = False
    last_hit = 0.0
    lufs_buf = np.zeros(int(48000 * 0.4), np.float32)
    lufs_i = 0
    drag = None
    last_y = 0
    press = None
    last_click = 0.0

    def click_ok() -> bool:
        nonlocal last_click
        now = time.time()
        if now - last_click < 0.16:
            return False
        last_click = now
        return True

    running = True
    while running:
        for ev in pygame.event.get():
            if ev.type == pygame.QUIT:
                running = False
            elif ev.type == pygame.KEYDOWN:
                if ev.key in (pygame.K_q, pygame.K_ESCAPE):
                    running = False
                elif ev.key == pygame.K_b and click_ok():
                    set_bypass(not ch.bypass)
                    persist()
                elif ev.key == pygame.K_g and click_ok():
                    ch.gain_match = not ch.gain_match
                    ch.apply()
                    persist()
                elif ev.key == pygame.K_s and click_ok():
                    ch.style = (ch.style + 1) % 3
                    ch.apply()
                    persist()
                elif ev.key == pygame.K_r and click_ok():
                    ch.apply_dict(DEFAULTS)
                    persist()
            elif ev.type == pygame.MOUSEBUTTONDOWN and ev.button == 1:
                if btn_bypass.collidepoint(ev.pos):
                    press = "bypass"
                elif btn_match.collidepoint(ev.pos):
                    press = "match"
                else:
                    hit = False
                    for i, r in enumerate(style_rects):
                        if r.collidepoint(ev.pos):
                            press = f"s{i}"
                            hit = True
                    if not hit:
                        for k in knobs:
                            if k.hit(ev.pos):
                                drag = k
                                last_y = ev.pos[1]
                                break
            elif ev.type == pygame.MOUSEBUTTONUP and ev.button == 1:
                if drag is not None:
                    persist()
                    drag = None
                elif press == "bypass" and btn_bypass.collidepoint(ev.pos) and click_ok():
                    set_bypass(not ch.bypass)
                    persist()
                elif press == "match" and btn_match.collidepoint(ev.pos) and click_ok():
                    ch.gain_match = not ch.gain_match
                    ch.apply()
                    persist()
                elif press and press.startswith("s") and click_ok():
                    i = int(press[1])
                    if style_rects[i].collidepoint(ev.pos):
                        ch.style = i
                        ch.apply()
                        persist()
                press = None
            elif ev.type == pygame.MOUSEWHEEL:
                pos = pygame.mouse.get_pos()
                for k in knobs:
                    if k.hit(pos):
                        k.nudge(0.04 * ev.y)
                        persist()
            elif ev.type == pygame.MOUSEMOTION and drag is not None:
                drag.nudge((last_y - ev.pos[1]) / 180.0)
                last_y = ev.pos[1]

        if live is not None:
            mono = live.snap()
            r = float(np.sqrt(np.mean(mono * mono)) + 1e-12)
            p = float(np.max(np.abs(mono)) + 1e-12)
            in_rms = out_rms = r
            in_peak = out_peak = p
            tp_hold = max(tp_hold * 0.999, db(p))
            n = len(mono)
            m = len(lufs_buf)
            if n >= m:
                lufs_buf[:] = mono[-m:]
                lufs_i = 0
            else:
                end = lufs_i + n
                if end <= m:
                    lufs_buf[lufs_i:end] = mono
                else:
                    k = m - lufs_i
                    lufs_buf[lufs_i:] = mono[:k]
                    lufs_buf[: n - k] = mono[k:]
                lufs_i = end % m
            loud = float(np.sqrt(np.mean(lufs_buf * lufs_buf)) + 1e-12)
            lufs = db(loud) - 0.7
            if db(r) > DETECT_DB:
                detected = True
                last_hit = time.time()
                spec = np.maximum(spectrum(mono, sr), spec * 0.86)
                hold = np.maximum(spec, hold * 0.992)
                take = min(len(mono), len(wave))
                wave = np.roll(wave, -take)
                wave[-take:] = mono[-take:]
                new_b = band_levels(mono, sr)
                bands = 0.40 * bands + 0.60 * new_b
                jump = bands - prev_bands
                if float(np.max(np.abs(jump))) > 0.16:
                    idx = int(np.argmax(np.abs(jump)))
                    flash = DJ_BANDS[idx][0]
                    flash_until = time.time() + 1.4
                prev_bands = bands.copy()
                dominant = int(np.argmax(bands)) if float(np.max(bands)) > 0.05 else dominant
                frozen_spec = spec.copy()
                frozen_bands = bands.copy()
            elif time.time() - last_hit > 0.35:
                detected = False
                spec = frozen_spec
                bands = frozen_bands
            if ch.gain_match and db(r) > DETECT_DB:
                target = -14.0
                ch.match_db = float(np.clip(target - db(r), -4.0, 4.0)) * 0.08 + ch.match_db * 0.92
                ch.apply()

        now = time.time()
        measuring = detected
        energy = float(np.clip((db(out_rms) + 48.0) / 36.0, 0.0, 1.0)) if measuring else 0.0
        show_flash = flash if measuring and now < flash_until else ""
        screen.blit(bg, (0, 0))
        pygame.draw.rect(screen, C_PANEL, (16, 16, 1368, 88), border_radius=16)
        pygame.draw.rect(screen, C_GREEN if measuring else C_ORANGE, (28, 32, 10, 52), border_radius=4)
        liquid_word(screen, pool, "LANDR", 48, 18, 46, energy, wave, now, measuring=measuring)
        liquid_word(screen, pool, "PRO MASTER", 220, 28, 30, energy * 0.85, wave[::-1] if wave.size else wave, now + 0.4, measuring=measuring)
        screen.blit(
            small.render(
                "15 LIVE FREQ  type + bands follow the set" if measuring else "15 STATIC  frequency not being measured",
                True, C_GREEN if measuring else C_ORANGE,
            ),
            (48, 74),
        )

        pygame.draw.rect(screen, C_TEAL if ch.gain_match else C_KNOB, btn_match, border_radius=10)
        screen.blit(font.render("13 GAIN MATCH", True, C_BG if ch.gain_match else C_WHITE), (btn_match.x + 8, btn_match.y + 10))
        pygame.draw.rect(screen, C_RED if ch.bypass else C_KNOB, btn_bypass, border_radius=10)
        screen.blit(font.render("14 BYPASS", True, C_WHITE), (btn_bypass.x + 16, btn_bypass.y + 10))
        badge = pygame.Rect(1188, 24, 184, 64)
        pygame.draw.rect(screen, C_GREEN if measuring else C_ORANGE, badge, border_radius=10)
        screen.blit(
            font.render("15 LIVE FREQ" if measuring else "15 STATIC", True, C_BG),
            (badge.x + 16, badge.y + 20),
        )

        pygame.draw.rect(screen, C_PANEL, (16, 116, 1368, 48), border_radius=12)
        screen.blit(font.render(f"17 IN  {inn}", True, C_SKY), (28, 130))
        screen.blit(font.render(f"18 OUT {out}", True, C_ORANGE), (520, 130))
        live_txt = "16 STREAM ON" if live is not None else "16 STREAM OFF"
        screen.blit(font.render(f"{live_txt}    R reset   G match   B bypass", True, C_MUTED), (900, 130))

        spec_rect = pygame.Rect(16, 176, 1120, 360)
        draw_liquid(
            screen, spec_rect, spec, wave, energy, measuring, bands, dominant,
            show_flash, pool, small, font,
        )

        draw_vumeter(screen, font, small, 1152, 176, 110, 360, db(out_rms), lufs, "20 LUFS", C_TEAL)
        draw_vumeter(screen, font, small, 1274, 176, 110, 360, db(out_peak), tp_hold, "21 TRUE PEAK", C_ORANGE)
        liquid_word(screen, pool, f"{lufs:5.1f}", 1158, 492, int(28 + 10 * energy), energy, wave, now, measuring=measuring)
        screen.blit(small.render("LUFS short", True, C_MUTED), (1162, 530))
        screen.blit(font.render(f"IN {db(in_peak):+5.1f}", True, C_SKY), (1278, 500))
        screen.blit(font.render(f"OUT {db(out_peak):+5.1f}", True, C_ORANGE), (1278, 522))

        screen.blit(small.render("02 STYLE", True, C_MUTED), (24, 592))
        for i, name in enumerate(STYLES):
            on = ch.style == i
            pygame.draw.rect(screen, C_TEAL if on else C_KNOB, style_rects[i], border_radius=10)
            if on:
                liquid_word(screen, pool, name, style_rects[i].x + 8, style_rects[i].y + 6, 20, energy, wave, now, measuring=measuring)
            else:
                screen.blit(font.render(name, True, C_WHITE), (style_rects[i].x + 14, style_rects[i].y + 10))

        for k in knobs:
            k.draw(screen, font, small)

        screen.blit(
            small.render("djay → BlackHole 2ch → this → PA    drag knobs    G gain-match    B bypass    Q quit", True, C_MUTED),
            (24, 796),
        )
        if err:
            screen.blit(font.render(f"ERR {err[:100]}", True, C_RED), (700, 792))
        pygame.display.flip()
        clock.tick(40)

    persist()
    if live is not None:
        live.close()
    pygame.quit()


if __name__ == "__main__":
    try:
        main()
    except KeyboardInterrupt:
        sys.exit(0)
