# tools/automaster_app/pipeline/reference_engine.py
"""Build per-genre reference profiles from commercial tracks in refs/<genre>/.

Profiles are the single source of truth for match_eq targets, width targets,
LUFS targets and the Phase C verification gate.
"""
import os, json, glob
import numpy as np
import soundfile as sf
import pyloudnorm as pyln

# ISO 1/3-octave centers 25 Hz .. 20 kHz
THIRD_OCTAVE_HZ = [25, 31.5, 40, 50, 63, 80, 100, 125, 160, 200, 250, 315,
                   400, 500, 630, 800, 1000, 1250, 1600, 2000, 2500, 3150,
                   4000, 5000, 6300, 8000, 10000, 12500, 16000, 20000]
WIDTH_BANDS_HZ = [250, 4000]
DEFAULT_PROFILE_DIR = os.path.join(os.path.dirname(__file__), "..", "references")


def third_octave_spectrum_db(audio, sr):
    """Per-band mean power in dB relative to overall, Welch-averaged.

    Welch (16k Hann segments, 50% overlap) instead of one full-length windowed
    FFT: a single Hann over a whole track tapers the first and last sections
    to silence, biasing the estimate toward the middle of the arrangement.
    """
    from scipy.signal import welch
    mono = audio.mean(axis=1) if audio.ndim == 2 else audio
    nper = int(min(16384, len(mono)))
    f, psd = welch(mono, fs=sr, nperseg=nper, noverlap=nper // 2)
    total = np.mean(psd) + 1e-24
    out = []
    for c in THIRD_OCTAVE_HZ:
        lo, hi = c / 2 ** (1 / 6), c * 2 ** (1 / 6)
        sel = psd[(f >= lo) & (f < hi)]
        band = np.mean(sel) if len(sel) else 0.0
        out.append(10 * np.log10(band / total + 1e-12))
    return out


def band_width_ratio(audio, sr):
    """S/(M+S) energy ratio per band (low/mid/high split at WIDTH_BANDS_HZ)."""
    from .filters import lr4_bands
    if audio.shape[1] == 1:
        return {name: 0.0 for name in ["low", "mid", "high"]}
    mid = (audio[:, 0] + audio[:, 1]) * 0.5
    side = (audio[:, 0] - audio[:, 1]) * 0.5
    names = ["low", "mid", "high"]
    ms = np.stack([mid, side], axis=1)
    bands = lr4_bands(ms, [float(b) for b in WIDTH_BANDS_HZ], sr)
    out = {}
    for name, b in zip(names, bands):
        em, es = np.mean(b[:, 0] ** 2), np.mean(b[:, 1] ** 2)
        out[name] = float(es / (em + es + 1e-12))
    return out


def build_profile(wav_paths, genre):
    from .tp_limiter import true_peak_db
    specs, lufss, lras, plrs, tps, widths = [], [], [], [], [], []
    for p in wav_paths:
        audio, sr = sf.read(p, always_2d=True)
        meter = pyln.Meter(sr)
        lufs = meter.integrated_loudness(audio)
        peak_db = 20 * np.log10(np.max(np.abs(audio)) + 1e-12)
        # loudness-normalize to -14 LUFS before spectral analysis
        norm = audio * 10 ** ((-14.0 - lufs) / 20.0)
        specs.append(third_octave_spectrum_db(norm, sr))
        lufss.append(lufs)
        plrs.append(peak_db - lufs)
        tps.append(true_peak_db(audio, sr))
        widths.append(band_width_ratio(audio, sr))
        # LRA: short-term loudness percentile spread
        st = _short_term_lufs(audio, sr, meter)
        lras.append(float(np.percentile(st, 95) - np.percentile(st, 10)) if len(st) else 0.0)
    return {
        "genre": genre,
        "n_refs": len(wav_paths),
        "lufs": float(np.mean(lufss)),
        "lra": float(np.mean(lras)),
        "plr": float(np.mean(plrs)),
        "true_peak_dbtp": float(np.mean(tps)),
        "third_octave_hz": [float(h) for h in THIRD_OCTAVE_HZ],
        "third_octave_db": [float(v) for v in np.mean(specs, axis=0)],
        "band_width_ratio": {k: float(np.mean([w[k] for w in widths]))
                             for k in ["low", "mid", "high"]},
        "width_bands_hz": WIDTH_BANDS_HZ,
    }


def _short_term_lufs(audio, sr, meter, win_s=3.0, hop_s=1.0):
    vals = []
    win, hop = int(win_s * sr), int(hop_s * sr)
    for i in range(0, len(audio) - win, hop):
        try:
            v = meter.integrated_loudness(audio[i:i + win])
            if np.isfinite(v):
                vals.append(v)
        except Exception:
            pass
    return np.array(vals)


def build_reference_profiles(refs_dir, out_dir=DEFAULT_PROFILE_DIR):
    os.makedirs(out_dir, exist_ok=True)
    built = []
    for gdir in sorted(glob.glob(os.path.join(refs_dir, "*"))):
        if not os.path.isdir(gdir):
            continue
        wavs = sorted(glob.glob(os.path.join(gdir, "*.wav")) +
                      glob.glob(os.path.join(gdir, "*.flac")))
        if not wavs:
            continue
        genre = os.path.basename(gdir)
        prof = build_profile(wavs, genre)
        with open(os.path.join(out_dir, f"{genre}.json"), "w") as fh:
            json.dump(prof, fh, indent=1)
        built.append(genre)
    return built


def load_profile(genre, profile_dir=DEFAULT_PROFILE_DIR):
    path = os.path.join(profile_dir, f"{genre}.json")
    if not os.path.exists(path):
        return None
    with open(path) as fh:
        return json.load(fh)


def _genre_key(name):
    return name.lower().replace(" ", "_").replace("-", "_")


def _tilted_third_octave_db(slope_db_per_oct=-3.0, ref_hz=1000.0):
    """⅓-octave spectrum dB relative to overall RMS for a tilted pink reference."""
    hz = np.array(THIRD_OCTAVE_HZ, dtype=np.float64)
    tilt = slope_db_per_oct * np.log2(np.maximum(hz, 20.0) / ref_hz)
    tilt -= np.mean(tilt)
    return [float(v) for v in tilt]


def _sibilance_from_third_octave(third_db):
    hz = np.array(THIRD_OCTAVE_HZ)
    db = np.array(third_db)
    e_sib = 10 ** (np.mean(db[(hz >= 4000) & (hz <= 8000)]) / 10.0)
    e_mid = 10 ** (np.mean(db[(hz >= 1000) & (hz <= 4000)]) / 10.0)
    return float(e_sib / (e_mid + 1e-12))


def builtin_profile(genre, genre_dna):
    """Fallback profile from genre DNA when no commercial refs exist (spec §3.1)."""
    dna = genre_dna or {}
    slope = float(dna.get("slope_target", dna.get("eq_emphasis", -3.0)))
    lufs = float(dna.get("target_lufs", dna.get("loudness_target", -10.0)))
    third_db = _tilted_third_octave_db(slope)
    return {
        "genre": genre,
        "n_refs": 0,
        "builtin": True,
        "lufs": lufs,
        "lra": 4.0,
        "plr": 8.0,
        "true_peak_dbtp": -1.0,
        "third_octave_hz": [float(h) for h in THIRD_OCTAVE_HZ],
        "third_octave_db": third_db,
        "band_width_ratio": {"low": 0.05, "mid": 0.3, "high": 0.45},
        "width_bands_hz": WIDTH_BANDS_HZ,
        "sibilance_ratio": _sibilance_from_third_octave(third_db),
    }


def ensure_builtin_profiles(profile_dir=DEFAULT_PROFILE_DIR, genres=None):
    """Write builtin JSON for genres missing a real-refs profile."""
    from ..genres import list_available_genres, get_genre_profile

    os.makedirs(profile_dir, exist_ok=True)
    built = []
    names = list(genres or [])
    for genre_name in names:
        key = _genre_key(genre_name) if isinstance(genre_name, str) else genre_name
        path = os.path.join(profile_dir, f"{key}.json")
        if os.path.exists(path):
            continue
        dna = get_genre_profile(genre_name.replace("_", " ").title()) or get_genre_profile("Tech House")
        prof = builtin_profile(key, dna)
        with open(path, "w") as fh:
            json.dump(prof, fh, indent=1)
        built.append(key)
    return built


def ensure_profile_for_genre(genre, profile_dir=DEFAULT_PROFILE_DIR, refs_dir=None):
    """Load, build-from-refs, or synthesize a builtin profile for one genre."""
    from ..genres import get_genre_profile

    key = _genre_key(genre)
    existing = load_profile(key, profile_dir)
    if existing:
        return existing
    os.makedirs(profile_dir, exist_ok=True)
    if refs_dir:
        gdir = os.path.join(refs_dir, key)
        if not os.path.isdir(gdir):
            for name in os.listdir(refs_dir):
                if _genre_key(name) == key:
                    gdir = os.path.join(refs_dir, name)
                    break
        if os.path.isdir(gdir):
            wavs = sorted(glob.glob(os.path.join(gdir, "*.wav")) +
                          glob.glob(os.path.join(gdir, "*.flac")))
            if wavs:
                prof = build_profile(wavs, key)
                with open(os.path.join(profile_dir, f"{key}.json"), "w") as fh:
                    json.dump(prof, fh, indent=1)
                return prof
    dna = get_genre_profile(genre.replace("_", " ").title()) or get_genre_profile("Tech House")
    prof = builtin_profile(key, dna)
    with open(os.path.join(profile_dir, f"{key}.json"), "w") as fh:
        json.dump(prof, fh, indent=1)
    return prof


def ensure_profiles(refs_dir=None, profile_dir=DEFAULT_PROFILE_DIR, genres=None):
    """Build profiles from refs/<genre>/ when present; builtin fallback per genre."""
    os.makedirs(profile_dir, exist_ok=True)
    built_real = []
    if refs_dir and os.path.isdir(refs_dir):
        built_real = build_reference_profiles(refs_dir, profile_dir)
    built_builtin = ensure_builtin_profiles(profile_dir, genres=genres)
    return {"real": built_real, "builtin": built_builtin}
