#!/usr/bin/env python3
"""LANDR Live — DDJ insert. LANDR Mastering Pro when installed; pedalboard twin fallback."""

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,
    Gain,
    HighpassFilter,
    HighShelfFilter,
    Limiter,
    LowShelfFilter,
    PeakFilter,
    Pedalboard,
)

ROOT = Path(__file__).resolve().parent
STATE_PATH = ROOT / "state_native.json"
WAVE_STATE_PATH = ROOT / "wave_persistent.json"
WAVE_LEN = 512
sys.path.insert(0, str(ROOT))
from pro_engine import ProSlot, get_slot, mount_chain  # noqa: E402

# DDJ-FLX USB only (djay OUT = controller). BlackHole = fallback if no controller.
WANT_IN = ("FLX", "DDJ-FLX", "DDJ", "Pioneer")
WANT_IN_FALLBACK = ("BlackHole 2ch", "BlackHole")
FORBID_IN = ("BlackHole 16ch", "BlackHole 64ch", "BlackHole 256ch")
WANT_OUT = (
    "External Headphones",
    "Headphones",
    "MacBook Air Speakers",
    "MacBook Pro Speakers",
    "Mac mini Speakers",
    "Built-in Output",
    "MacBook 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 = -52.0
LIVE_SR = 48000
# 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)

# Bang Wong accents — unique per numbered control (colorblind-safe)
WONG = (
    (0, 158, 115),
    (213, 94, 0),
    (0, 114, 178),
    (204, 121, 167),
    (240, 228, 66),
    (86, 180, 233),
    (230, 159, 0),
)
WONG_LOCK = WONG[0]
WONG_WAIT = WONG[1]
WONG_IN = WONG[5]
WONG_OUT = WONG[6]
WONG_MASTER = WONG[2]


def wong_for_num(num: str) -> tuple[int, int, int]:
    digits = "".join(c for c in num if c.isdigit())
    return WONG[int(digits or "0") % len(WONG)]


def draw_route_card(
    surf, small, font, rect: pygame.Rect, num: str, title: str, detail: str, accent: tuple[int, int, int],
) -> None:
    pygame.draw.rect(surf, (14, 10, 34), rect, border_radius=14)
    pygame.draw.rect(surf, accent, (rect.x + 2, rect.y + 10, 4, rect.h - 20), border_radius=2)
    surf.blit(small.render(num, True, accent), (rect.x + 18, rect.y + 12))
    surf.blit(font.render(title, True, C_WHITE), (rect.x + 18, rect.y + 32))
    surf.blit(small.render(detail[:40], True, C_MUTED), (rect.x + 18, rect.y + 56))


def draw_key_chip(surf, small, x: int, y: int, text: str) -> int:
    chip = small.render(text, True, C_SKY)
    r = pygame.Rect(x, y, chip.get_width() + 18, chip.get_height() + 10)
    pygame.draw.rect(surf, C_WELL, r, border_radius=8)
    surf.blit(chip, (x + 9, y + 5))
    return r.right + 10

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,
}

# DJ insert — unity-ish; no comp/loudness lift on noise floor (Pioneer IN → External OUT)
CLEAN = {
    "in_gain": 0.0,
    "style": 1,
    "eq_low": 0.0,
    "eq_mid": 0.0,
    "eq_high": 0.0,
    "presence": 0.0,
    "deess": 0.0,
    "width": 0.5,
    "comp": 0.0,
    "character": 0.0,
    "sat": 0.0,
    "loudness": 0.0,
    "bypass": True,
    "gain_match": False,
}


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


def sd_list(kind: str, forbid: tuple[str, ...] = ()) -> list[tuple[int, str]]:
    out: list[tuple[int, str]] = []
    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
        out.append((i, name))
    return out


def sd_pick(kind: str, wants: tuple[str, ...], forbid: tuple[str, ...] = ()) -> tuple[int, str]:
    found = sd_list(kind, forbid)
    for want in wants:
        for i, name in found:
            if want.lower() in name.lower():
                return i, name
    if found:
        return found[0]
    sys.exit("No audio devices.")


def is_ddj_in(name: str) -> bool:
    low = name.lower()
    return any(x in low for x in ("ddj", "flx", "pioneer", "pioneer dj"))


def sd_pick_live_in() -> tuple[int, str]:
    """DDJ-FLX by name — never BlackHole if controller is plugged in."""
    found = sd_list("in", FORBID_IN)
    for want in WANT_IN:
        for i, name in found:
            if want.lower() in name.lower():
                return i, name
    for want in WANT_IN_FALLBACK:
        for i, name in found:
            if want.lower() in name.lower():
                return i, name
    return sd_pick("in", WANT_IN + WANT_IN_FALLBACK, forbid=FORBID_IN)


def pick_ddj_stereo(indata: np.ndarray, pair_hint: int = 0) -> tuple[np.ndarray, int]:
    """DDJ-FLX: djay FLX = USB ch 1–2, External/booth often ch 3–4."""
    x = np.nan_to_num(np.asarray(indata, dtype=np.float32), copy=False)
    if x.ndim == 1:
        x = np.column_stack([x, x])
    ncol = x.shape[1]
    if ncol <= 2:
        return np.clip(x[:, :2], -1.0, 1.0), 0
    best_off = 0
    best_e = -1.0
    for off in range(0, min(ncol - 1, 6), 2):
        pair = x[:, off : off + 2]
        e = float(np.sqrt(np.mean(pair * pair) + 1e-18))
        if e > best_e:
            best_e = e
            best_off = off
    if pair_hint >= 0 and pair_hint + 1 < ncol:
        hint = x[:, pair_hint : pair_hint + 2]
        e_hint = float(np.sqrt(np.mean(hint * hint) + 1e-18))
        if e_hint >= best_e * 0.85:
            best_off = pair_hint
    stereo = x[:, best_off : best_off + 2]
    return np.clip(stereo, -1.0, 1.0), best_off


def sd_cycle(kind: str, cur: int, step: int, forbid: tuple[str, ...] = ()) -> tuple[int, str]:
    found = sd_list(kind, forbid)
    if not found:
        sys.exit("No audio devices.")
    ids = [i for i, _ in found]
    try:
        pos = ids.index(cur)
    except ValueError:
        pos = 0
    pos = (pos + step) % len(found)
    return found[pos]


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=30.0)
        self.comp_p = Compressor(threshold_db=-18.0, ratio=2.2, attack_ms=10.0, release_ms=120.0)
        self.low = LowShelfFilter(cutoff_frequency_hz=80.0, gain_db=0.4)
        self.mid = PeakFilter(cutoff_frequency_hz=350.0, gain_db=0.0, q=0.60)
        self.pres = PeakFilter(cutoff_frequency_hz=2500.0, gain_db=-0.6, q=0.85)
        self.deess_p = PeakFilter(cutoff_frequency_hz=7000.0, gain_db=0.0, q=1.20)
        self.high = HighShelfFilter(cutoff_frequency_hz=8500.0, gain_db=0.3)
        self.air = HighShelfFilter(cutoff_frequency_hz=11000.0, gain_db=0.2)
        self.g_loud = Gain(gain_db=1.2)
        self.g_match = Gain(gain_db=0.0)
        self.lim = Limiter(threshold_db=-1.8, release_ms=90.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.g_loud, self.g_match, self.lim,
        ]
        self.board = Pedalboard(self.parts)
        self.pro: ProSlot | None = None
        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))
        loud_n = float(np.clip(self.loudness, 0.0, 1.0))
        loud = 0.0 if loud_n < 0.04 else (-2.0 + 6.0 * loud_n)
        attack = 3.0 + 18.0 * ch
        release = 40.0 + 179.0 * ch
        low = self.eq_low
        mid = self.eq_mid
        high = self.eq_high
        master_on = c >= 0.04 or loud_n >= 0.04 or ch >= 0.04
        if master_on:
            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
        if c < 0.04:
            self.comp_p.threshold_db = 0.0
            self.comp_p.ratio = 1.0
            self.comp_p.attack_ms = 5.0
            self.comp_p.release_ms = 50.0
        else:
            self.comp_p.threshold_db = -6.0 - 14.0 * c
            self.comp_p.ratio = 1.4 + 2.6 * 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 = 0.15 + (w - 0.5) * 0.8
        self.g_loud.gain_db = loud
        self.g_match.gain_db = float(np.clip(self.match_db, -3.0, 3.0)) if self.gain_match else 0.0
        self.lim.threshold_db = -2.0 if loud_n < 0.04 and c < 0.04 else -1.8

    def process(self, block: np.ndarray, sr: float) -> np.ndarray:
        if self.bypass:
            return np.clip(block, -1.0, 1.0)
        if self.pro is not None:
            pro_out = self.pro.process(block, sr)
            if pro_out is not None:
                return pro_out
        x = np.ascontiguousarray(block.T)
        y = self.board(x, float(sr))
        if y.ndim == 1:
            y = np.stack([y, y], axis=0)
        y = np.nan_to_num(y, copy=False)
        w = float(np.clip(self.width, 0.0, 1.0))
        s = float(np.clip(self.sat, 0.0, 1.0))
        if abs(w - 0.5) >= 0.03 or s >= 0.02:
            mid = 0.5 * (y[0] + y[1])
            side = 0.5 * (y[0] - y[1])
            side *= 0.90 + 0.20 * w
            if s > 0.02:
                mid = np.tanh(mid * (1.0 + 0.55 * s)) / (1.0 + 0.15 * s)
            y = np.stack([mid + side, mid - side], axis=0)
        return np.clip(np.ascontiguousarray(y.T), -1.0, 1.0)


class Live:
    """Single duplex stream — same clock in/out (no ring drift or block resample)."""

    def __init__(self, ch: Chain, in_dev: int, out_dev: int, sr: int) -> None:
        self.ch = ch
        self.lock = threading.Lock()
        self.buf = np.zeros(FFT_N, np.float32)
        self.err = ""
        self.in_dev = in_dev
        self.out_dev = out_dev
        in_info = sd.query_devices(in_dev)
        out_info = sd.query_devices(out_dev)
        self.sr = LIVE_SR
        self.in_ch = 2
        self.active_pair = 0
        in_name = str(in_info.get("name", ""))
        max_in = int(in_info.get("max_input_channels") or 2)
        if is_ddj_in(in_name) and max_in >= 4:
            self.in_ch = min(4, max_in)
        block = 256
        base_kw = dict(blocksize=block, dtype="float32", latency="low")

        def _callback(indata, outdata, frames, time_info, status) -> None:
            raw = np.nan_to_num(np.asarray(indata, dtype=np.float32), copy=False)
            xin, self.active_pair = pick_ddj_stereo(raw, self.active_pair)
            try:
                y = self.ch.process(xin, float(self.sr))
            except (ValueError, RuntimeError, OSError):
                y = xin
            y = np.asarray(y, dtype=np.float32)
            if y.ndim == 1:
                y = np.column_stack([y, y])
            y = np.clip(np.nan_to_num(y, copy=False), -1.0, 1.0)
            n = min(len(y), len(outdata))
            outdata[:n, :2] = y[:n, :2]
            if n < len(outdata):
                outdata[n:, :] = 0.0
            if raw.ndim > 1 and raw.shape[1] > 1:
                mono = np.max(np.abs(raw), axis=1)
            else:
                mono = xin.mean(axis=1)
            k = min(len(mono), FFT_N)
            with self.lock:
                self.buf[:-k] = self.buf[k:]
                self.buf[-k:] = mono[-k:]
        self.stream = None
        in_native = int(in_info.get("default_samplerate") or LIVE_SR)
        out_native = int(out_info.get("default_samplerate") or LIVE_SR)
        rates = []
        for r in (LIVE_SR, in_native, out_native, 44100):
            if r not in rates:
                rates.append(r)
        last_err = ""
        ch_tries = (self.in_ch, 2) if self.in_ch > 2 else (2,)
        for try_sr in rates:
            for in_ch in ch_tries:
                try:
                    kw = {**base_kw, "channels": (in_ch, 2)}
                    self.stream = sd.Stream(
                        device=(in_dev, out_dev),
                        samplerate=try_sr,
                        callback=_callback,
                        **kw,
                    )
                    self.sr = try_sr
                    self.in_ch = in_ch
                    break
                except Exception as e:
                    last_err = str(e)
                    self.stream = None
            if self.stream is not None:
                break
        if self.stream is None:
            raise RuntimeError(last_err or "duplex open failed")
        if "blackhole" in str(out_info.get("name", "")).lower():
            self.err = "Output cannot be BlackHole"
        self.stream.start()

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

    def close(self) -> None:
        st = getattr(self, "stream", None)
        if st is None:
            return
        try:
            st.stop()
            st.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)
        accent = wong_for_num(self.num)
        pygame.draw.line(
            surf, accent, (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 wave_has_shape(wave: np.ndarray) -> bool:
    return wave.size > 8 and float(np.max(np.abs(wave))) > 1e-5


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,
    persistent: bool = True,
) -> None:
    show = measuring or (persistent and wave_has_shape(wave))
    if not show:
        glyph = pool.get(base).render(text, True, C_MUTED)
        surf.blit(glyph, (x, y))
        return
    if not measuring:
        energy *= 0.45
    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 wave_has_shape(wave):
        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:
            alpha = 210 if live else 120
            pygame.draw.lines(well, (170, 150, 255, alpha), False, wpts, 2 if live else 1)
    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 SPECTRUM   {now}", True, WONG_LOCK), (rect.x + 20, rect.y + 10))
        if flash:
            surf.blit(font.render(f"23 {flash}", True, WONG[4]), (rect.x + 420, rect.y + 10))
    else:
        surf.blit(font.render("24 Waveform hold", True, WONG_WAIT), (rect.x + 20, rect.y + 10))
        surf.blit(small.render("last spectrum until new input", 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("LANDR Live · master insert")
    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))
    screen.blit(huge.render("LANDR Live", True, C_WHITE), (32, 36))
    screen.blit(small.render("Starting audio engine…", True, WONG_IN), (32, 88))
    pygame.display.flip()
    pygame.event.pump()

    in_dev, inn = sd_pick_live_in()
    out_dev, out = sd_pick("out", WANT_OUT, forbid=FORBID_OUT)
    inn_low = inn.lower()
    if is_ddj_in(inn) and "16ch" not in inn_low:
        err_boot = ""
    elif "blackhole" in inn_low and "2ch" in inn_low:
        err_boot = ""
    else:
        err_boot = f"17 Input: connect DDJ-FLX USB (now: {inn[:44]})"
    if "blackhole" in out.lower():
        err_boot = (err_boot + " | " if err_boot else "") + f"OUT cannot be BlackHole (got {out})"
    ch = Chain()
    mount_chain(ch)
    use_clean = "--clean" in sys.argv
    if use_clean:
        ch.apply_dict(CLEAN)
        ch.gain_match = False
        ch.match_db = 0.0
        if ch.pro is not None and ch.pro.using_pro():
            ch.bypass = False
        ch.apply()
        STATE_PATH.write_text(json.dumps(ch.dump(), indent=2))
    elif STATE_PATH.exists():
        try:
            saved = json.loads(STATE_PATH.read_text())
            ch.apply_dict(saved)
            if saved.get("gain_match"):
                ch.gain_match = False
                ch.apply()
        except (OSError, json.JSONDecodeError, TypeError, ValueError):
            pass
    sr = LIVE_SR
    live = None
    err = err_boot
    try:
        live = Live(ch, in_dev, out_dev, sr)
    except (OSError, RuntimeError, ValueError) as e:
        err = (err + " | " if err else "") + str(e)

    def persist() -> None:
        STATE_PATH.write_text(json.dumps(ch.dump(), indent=2))
        try:
            WAVE_STATE_PATH.write_text(
                json.dumps(
                    {
                        "wave": wave.tolist(),
                        "persist_energy": persist_energy,
                        "frozen_spec": frozen_spec.tolist(),
                        "frozen_bands": frozen_bands.tolist(),
                        "dominant": dominant,
                    }
                ),
                encoding="utf-8",
            )
        except OSError:
            pass

    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(WAVE_LEN, np.float32)
    persist_energy = 0.0
    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
    if WAVE_STATE_PATH.is_file():
        try:
            ws = json.loads(WAVE_STATE_PATH.read_text(encoding="utf-8"))
            w = np.asarray(ws.get("wave", []), dtype=np.float32)
            if w.size == WAVE_LEN:
                wave = w
            persist_energy = float(ws.get("persist_energy", 0.0))
            if wave_has_shape(wave) and persist_energy < 0.12:
                persist_energy = 0.28
            fs = np.asarray(ws.get("frozen_spec", []), dtype=np.float32)
            if fs.size == BARS:
                frozen_spec = fs
                spec = fs.copy()
            fb = np.asarray(ws.get("frozen_bands", []), dtype=np.float32)
            if fb.size == len(DJ_BANDS):
                frozen_bands = fb
                bands = fb.copy()
            dominant = int(ws.get("dominant", 0))
        except Exception:
            pass
    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
    in_lock_sent = False
    lufs_buf = np.zeros(int(48000 * 0.4), np.float32)
    lufs_i = 0
    drag = None
    last_y = 0
    press = None
    last_click = 0.0
    last_wave_save = 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.key == pygame.K_c and click_ok():
                    ch.apply_dict(CLEAN)
                    ch.gain_match = False
                    ch.apply()
                    persist()
                elif ev.key in (pygame.K_1, pygame.K_2) and click_ok():
                    step = -1 if ev.key == pygame.K_1 else 1
                    in_dev, inn = sd_cycle("in", in_dev, step, FORBID_IN)
                    if live is not None:
                        live.close()
                    try:
                        live = Live(ch, in_dev, out_dev, 0)
                        err = ""
                    except Exception as e:
                        live = None
                        err = str(e)
                elif ev.key in (pygame.K_3, pygame.K_4) and click_ok():
                    step = -1 if ev.key == pygame.K_3 else 1
                    out_dev, out = sd_cycle("out", out_dev, step, FORBID_OUT)
                    if live is not None:
                        live.close()
                    try:
                        live = Live(ch, in_dev, out_dev, 0)
                        err = ""
                    except Exception as e:
                        live = None
                        err = str(e)
            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 max(db(r), db(p)) > DETECT_DB:
                if not detected and not in_lock_sent:
                    in_lock_sent = True
                    slot = get_slot()

                    def _signal_lock() -> None:
                        ProSlot.notify("LANDR Live", "Signal lock")
                        if slot.using_pro():
                            slot.open_editor()

                    threading.Thread(target=_signal_lock, daemon=True).start()
                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
                wave *= 0.9990
                persist_energy *= 0.993
            if ch.gain_match and db(r) > -32.0:
                target = -14.0
                ch.match_db = float(np.clip(target - db(r), -2.0, 2.0)) * 0.03 + ch.match_db * 0.97
                ch.g_match.gain_db = float(np.clip(ch.match_db, -2.0, 2.0))

        now = time.time()
        measuring = detected
        energy = float(np.clip((db(out_rms) + 48.0) / 36.0, 0.0, 1.0))
        if measuring:
            persist_energy = float(np.clip(0.65 * persist_energy + 0.35 * energy, 0.0, 1.0))
        display_energy = energy if measuring else persist_energy
        holding_wave = (not measuring) and wave_has_shape(wave)
        show_flash = flash if measuring and now < flash_until else ""
        screen.blit(bg, (0, 0))
        pygame.draw.rect(screen, (14, 10, 34), (16, 16, 1368, 96), border_radius=16)
        screen.blit(big.render("LANDR Live", True, C_WHITE), (28, 24))
        screen.blit(small.render("Master insert · controller in → speakers out", True, WONG_IN), (28, 62))

        pygame.draw.rect(screen, C_TEAL if ch.gain_match else C_KNOB, btn_match, border_radius=10)
        screen.blit(font.render("15 Gain match", True, C_BG if ch.gain_match else C_WHITE), (btn_match.x + 6, 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("16 Bypass", True, C_WHITE), (btn_bypass.x + 10, btn_bypass.y + 10))

        if measuring:
            pill_txt, pill_col = "Signal lock", WONG_LOCK
        elif holding_wave:
            pill_txt, pill_col = "Waveform hold", WONG_WAIT
        else:
            pill_txt, pill_col = "No input signal", WONG_WAIT
        pill = pygame.Rect(960, 28, 424, 44)
        pygame.draw.rect(screen, pill_col, pill, border_radius=10)
        screen.blit(font.render(pill_txt, True, (12, 10, 24)), (pill.x + 16, pill.y + 10))
        stream_txt = "25 Engine on" if live is not None else "25 Engine off"
        screen.blit(small.render(stream_txt, True, WONG_MASTER if live else C_MUTED), (pill.x + 16, pill.y + 26))

        in_title = "Input"
        pair_note = ""
        if live is not None and is_ddj_in(inn):
            p = int(getattr(live, "active_pair", 0))
            if p > 0:
                pair_note = f" · USB {p + 1}–{p + 2}"
        if ch.bypass:
            proc_title, proc_detail = "Process", "Passthrough"
        elif ch.pro is not None:
            proc_title, proc_detail = "Process", ch.pro.status
        else:
            proc_title, proc_detail = "Process", "Twin pedalboard"
        draw_route_card(
            screen, small, font, pygame.Rect(16, 120, 448, 76),
            "17", in_title, f"{inn}{pair_note}", WONG_IN,
        )
        draw_route_card(
            screen, small, font, pygame.Rect(476, 120, 448, 76),
            "19", proc_title, proc_detail, WONG_MASTER,
        )
        draw_route_card(
            screen, small, font, pygame.Rect(936, 120, 448, 76),
            "18", "Output", out[:40], WONG_OUT,
        )

        spec_rect = pygame.Rect(16, 208, 1120, 328)
        draw_liquid(
            screen, spec_rect, spec, wave, display_energy, measuring, bands, dominant,
            show_flash, pool, small, font,
        )

        draw_vumeter(screen, font, small, 1152, 208, 110, 328, db(out_rms), lufs, "20 LUFS (short)", WONG[2])
        draw_vumeter(screen, font, small, 1274, 208, 110, 328, db(out_peak), tp_hold, "21 True peak", WONG[1])
        screen.blit(font.render(f"{lufs:5.1f}", True, C_WHITE), (1158, 520))
        screen.blit(small.render("LUFS", True, C_MUTED), (1162, 548))
        screen.blit(font.render(f"17 In {db(in_peak):+5.1f} dBFS", True, C_SKY), (1274, 500))
        screen.blit(font.render(f"18 Out {db(out_peak):+5.1f} dBFS", True, C_ORANGE), (1274, 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)
            col = wong_for_num(f"0{i + 2}")
            surf_txt = font.render(name, True, (12, 10, 24) if on else C_WHITE)
            screen.blit(surf_txt, (style_rects[i].x + 14, style_rects[i].y + 10))

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

        chip_x = 24
        chip_x = draw_key_chip(screen, small, chip_x, 796, "1/2 Input")
        chip_x = draw_key_chip(screen, small, chip_x, 796, "3/4 Output")
        chip_x = draw_key_chip(screen, small, chip_x, 796, "G Match")
        chip_x = draw_key_chip(screen, small, chip_x, 796, "B Bypass")
        chip_x = draw_key_chip(screen, small, chip_x, 796, "Q Quit")
        screen.blit(small.render("djay output = DDJ-FLX", True, C_MUTED), (chip_x + 8, 800))
        if err:
            screen.blit(font.render(f"ERR {err[:100]}", True, C_RED), (700, 792))
        if (measuring or holding_wave) and now - last_wave_save > 2.0:
            persist()
            last_wave_save = now
        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)
    except Exception:
        import traceback

        crash = ROOT / "LIVE_CRASH.log"
        crash.write_text(traceback.format_exc(), encoding="utf-8")
        traceback.print_exc()
        sys.exit(1)
