#!/usr/bin/env python3
"""Integration gate: master library tracks and score via Phase C verification."""
import argparse
import csv
import json
import os
import re
import resource
import sys
import time
from pathlib import Path

import numpy as np
import soundfile as sf
from rich.console import Console
from rich.table import Table

sys.path.insert(0, os.path.dirname(__file__))

from automaster_app.genres import get_genre_profile, list_available_genres
from automaster_app.pipeline.orchestrator import master_track
from automaster_app.pipeline.reference_engine import (
    DEFAULT_PROFILE_DIR,
    ensure_profile_for_genre,
    ensure_profiles,
    load_profile,
)
from automaster_app.pipeline.verify import export_ab_snippet

console = Console()
RESULTS_DIR = os.path.join("tmp", "evaluation")
RESULTS_CSV = os.path.join(RESULTS_DIR, "results.csv")
DEFAULT_LIBRARY = "/home/user/Downloads/2026"
SR = 44100
SKIP_DIRS = {"Losers", "Mastered"}
SOURCE_SKIP_RE = re.compile(
    r"MASTERED|_master\.|_AB\.|_Storyboard\.|_v\d+_master",
    re.IGNORECASE,
)


def get_track_root_name(filename):
    name = filename.replace(" MASTERED", "").replace(".wav", "").replace(".mp3", "")
    name = re.sub(r"_(v|AC)\d+.*", "", name)
    name = re.sub(r"_\d+_\d+", "", name)
    return name.replace("(Cover)", "").strip()


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


def infer_genre(path, default="other"):
    stem = Path(path).stem.lower()
    for genre_name in sorted(list_available_genres(), key=lambda g: len(g), reverse=True):
        key = _genre_key(genre_name)
        if key in stem or genre_name.lower() in stem:
            return key
    return default


def discover_source_tracks(root_path, limit=None):
    tracks = []
    for root, dirs, files in os.walk(root_path):
        dirs[:] = [d for d in dirs if d not in SKIP_DIRS]
        if any(skip in root for skip in SKIP_DIRS):
            continue
        for file in sorted(files):
            low = file.lower()
            if not low.endswith((".wav", ".flac", ".mp3")):
                continue
            if SOURCE_SKIP_RE.search(file):
                continue
            tracks.append(os.path.join(root, file))
            if limit and len(tracks) >= limit:
                return tracks
    return tracks


def _note_freq(note, octave=3):
    names = ["C", "C#", "D", "D#", "E", "F", "F#", "G", "G#", "A", "A#", "B"]
    base = 16.35 * (2 ** octave)
    idx = names.index(note)
    return base * (2 ** (idx / 12.0))


def _pink_bed(n, ch, seed, slope_db_per_oct=-3.0):
    rng = np.random.default_rng(seed)
    white = rng.standard_normal((n, ch))
    f = np.fft.rfft(white, axis=0)
    freqs = np.fft.rfftfreq(n, 1 / SR)
    f[1:] /= np.sqrt(np.arange(1, f.shape[0]))[:, None]
    shape = (np.maximum(freqs, 20.0) / 1000.0) ** (slope_db_per_oct / 6.020)
    f *= shape[:, None]
    bed = np.fft.irfft(f, n=n, axis=0)
    return bed / (np.max(np.abs(bed)) + 1e-12)


def _bandlimit(audio, cutoff_hz=15000.0):
    f = np.fft.rfft(audio, axis=0)
    freqs = np.fft.rfftfreq(audio.shape[0], 1 / SR)
    f[freqs > cutoff_hz] = 0
    return np.fft.irfft(f, n=audio.shape[0], axis=0)


def _shape_width(audio, band_width_ratio):
    """Rough M/S band scaling toward builtin width targets."""
    from automaster_app.pipeline.filters import lr4_bands

    targets = band_width_ratio
    splits = [250.0, 4000.0]
    mid = (audio[:, 0] + audio[:, 1]) * 0.5
    side = (audio[:, 0] - audio[:, 1]) * 0.5
    ms = np.stack([mid, side], axis=1)
    bands = lr4_bands(ms, splits, SR)
    out_m = np.zeros_like(mid)
    out_s = np.zeros_like(side)
    for name, band in zip(["low", "mid", "high"], bands):
        em = np.mean(band[:, 0] ** 2) + 1e-12
        es = np.mean(band[:, 1] ** 2) + 1e-12
        cur = es / (em + es)
        tgt = float(targets[name])
        want = tgt / max(1.0 - tgt, 1e-6) * em
        scale = float(np.sqrt(want / (es + 1e-12)))
        out_m += band[:, 0]
        out_s += band[:, 1] * scale
    return np.stack([out_m + out_s, out_m - out_s], axis=1)


def _align_spectrum(audio, prof, rolloff_hz=20000.0, max_passes=6, target_dev=1.5):
    from automaster_app.pipeline.match_eq import apply_match_eq
    from automaster_app.pipeline.reference_engine import third_octave_spectrum_db
    import pyloudnorm as pyln

    z = audio.copy()
    meter = pyln.Meter(SR)
    hz = np.array(prof["third_octave_hz"])
    ref = np.array(prof["third_octave_db"])
    sel = (hz >= 50) & (hz <= 16000)
    for _ in range(max_passes):
        z = apply_match_eq(z, SR, prof, pre_extension_rolloff_hz=rolloff_hz)
        lufs = meter.integrated_loudness(z)
        zn = z * 10 ** ((-14.0 - lufs) / 20.0)
        spec = np.array(third_octave_spectrum_db(zn, SR))
        dev = float(np.mean(np.abs(spec[sel] - ref[sel])))
        if dev <= target_dev:
            break
    return z


def _musical_layers(n, t, bpm, chord_notes, rng, level=1.0):
    audio = np.zeros((n, 2), dtype=np.float64)
    bar_len = int(SR * 60.0 / bpm * 4)
    chord_freqs = [_note_freq(note, 3) for note in chord_notes]
    for bar_start in range(0, n, bar_len):
        bar_end = min(n, bar_start + bar_len)
        seg = slice(bar_start, bar_end)
        tt = t[seg] - t[bar_start]
        chord = sum(0.04 * level * np.sin(2 * np.pi * f * tt) for f in chord_freqs)
        audio[seg] += chord[:, None] * np.array([1.0, 0.94])
    kick_period = int(SR * 60.0 / bpm)
    for i in range(0, n, kick_period):
        env_len = min(kick_period // 2, n - i)
        env = np.exp(-np.arange(env_len) / (SR * 0.04))
        kick = 0.16 * level * np.sin(2 * np.pi * 55 * np.arange(env_len) / SR) * env
        audio[i:i + env_len] += kick[:, None]
    hat_period = kick_period // 2
    for i in range(hat_period, n, hat_period):
        burst = min(400, n - i)
        noise = rng.standard_normal(burst) * np.hanning(burst) * 0.025 * level
        audio[i:i + burst, 0] += noise
        audio[i:i + burst, 1] += noise * 0.9
    bass_root = chord_freqs[0] / 2.0
    audio += (0.07 * level * np.sin(2 * np.pi * bass_root * t))[:, None]
    return audio


def _synth_reference_and_source(seconds, bpm, chord_notes, genre, seed=0):
    """Build a profile-aligned reference master and a band-limited Suno-like source."""
    import pyloudnorm as pyln
    from automaster_app.pipeline.reference_engine import builtin_profile, load_profile
    from automaster_app.pipeline.tp_limiter import limit

    dna = get_genre_profile(genre.replace("_", " ").title()) or get_genre_profile("Tech House")
    prof = load_profile(genre) or builtin_profile(genre, dna)
    slope = float((dna or {}).get("slope_target", -3.0))
    rng = np.random.default_rng(seed)
    n = int(SR * seconds)
    t = np.arange(n) / SR

    bed = _pink_bed(n, 2, seed, slope) * 0.5
    bed += _musical_layers(n, t, bpm, chord_notes, rng, level=0.8)
    bed = _shape_width(bed, prof["band_width_ratio"])
    reference = _align_spectrum(bed, prof, rolloff_hz=20000.0)
    reference, _ = limit(reference, SR, target_lufs=float(prof["lufs"]), ceiling_dbtp=-1.0)

    source = _bandlimit(reference, 16500.0)
    source += rng.standard_normal((n, 2)) * 0.001
    meter = pyln.Meter(SR)
    src_lufs = meter.integrated_loudness(source)
    source = source * 10 ** ((-14.0 - src_lufs) / 20.0)
    peak = np.max(np.abs(source)) + 1e-12
    source = source / peak * 0.85
    return reference, source


def generate_synthetic_corpus(out_dir, refs_dir=None, force=False):
    os.makedirs(out_dir, exist_ok=True)
    specs = [
        ("corpus_tech_house.wav", "tech_house", 124, ["G", "A", "B", "A"], 16),
        ("corpus_modern_pop.wav", "modern_pop", 110, ["C", "G", "A", "F"], 17),
        ("corpus_melodic_techno.wav", "melodic_techno", 128, ["A", "F", "C", "G"], 18),
    ]
    paths = []
    for name, genre, bpm, chord, seed in specs:
        path = os.path.join(out_dir, name)
        if force or not os.path.exists(path):
            reference, source = _synth_reference_and_source(
                30, bpm, chord, genre, seed=seed)
            sf.write(path, source, SR)
            if refs_dir:
                gdir = os.path.join(refs_dir, genre)
                os.makedirs(gdir, exist_ok=True)
                sf.write(os.path.join(gdir, f"{genre}_ref.wav"), reference, SR)
        paths.append(path)
    return paths


def _peak_rss_mb():
    rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
    if sys.platform == "darwin":
        return rss / (1024 * 1024)
    return rss / 1024


def _failed_checks(verification):
    if verification.get("skipped"):
        return "no_profile"
    checks = verification.get("checks", {})
    return ",".join(
        k for k, v in checks.items()
        if not v.get("skipped") and not v.get("pass", False)
    )


def _reference_wav_for_genre(genre, refs_dir):
    if not refs_dir or not os.path.isdir(refs_dir):
        return None
    gdir = os.path.join(refs_dir, genre)
    if not os.path.isdir(gdir):
        for name in os.listdir(refs_dir):
            if _genre_key(name) == genre:
                gdir = os.path.join(refs_dir, name)
                break
    if not os.path.isdir(gdir):
        return None
    for ext in ("*.wav", "*.flac"):
        hits = sorted(Path(gdir).glob(ext))
        if hits:
            return str(hits[0])
    return None


def evaluate_track(path, genre, profile_dir, out_dir, refs_dir=None,
                   no_stems=False, light=False):
    wall0 = time.time()
    out_path, report = master_track(
        path,
        genre=genre,
        no_stems=no_stems,
        light=light,
        out_dir=out_dir,
        profile_dir=profile_dir,
    )
    wall_s = round(time.time() - wall0, 1)
    peak_rss_mb = round(_peak_rss_mb(), 1)

    verification = report.get("verification", {})
    if verification.get("skipped"):
        profile = load_profile(genre, profile_dir)
        if profile:
            from automaster_app.pipeline.verify import verify_master
            import soundfile as sf_read

            master_audio, sr = sf_read.read(out_path, always_2d=True)
            input_audio, _ = sf_read.read(path, always_2d=True)
            verification = verify_master(
                master_audio,
                sr,
                profile,
                input_audio=input_audio,
                separation_quality=report.get("phases", {}).get("A"),
            )
            report["verification"] = verification

    ref_wav = _reference_wav_for_genre(genre, refs_dir)
    if ref_wav:
        master_audio, sr = sf.read(out_path, always_2d=True)
        ref_audio, ref_sr = sf.read(ref_wav, always_2d=True)
        if ref_sr != sr:
            from scipy.signal import resample_poly
            ref_audio = resample_poly(ref_audio, sr, ref_sr, axis=0)
        ab_path = os.path.join(
            out_dir, f"{Path(path).stem}_AB.wav")
        export_ab_snippet(master_audio, ref_audio, sr, ab_path)

    checks = verification.get("checks", {})
    spec = checks.get("spectral_distance_db", {})
    row = {
        "file": path,
        "genre": genre,
        "phase_a": report.get("phases", {}).get("A", ""),
        "elapsed_s": report.get("elapsed_s", wall_s),
        "wall_clock_s": wall_s,
        "peak_rss_mb": peak_rss_mb,
        "output_lufs": report.get("output_lufs"),
        "true_peak_dbtp": report.get("output_true_peak_dbtp"),
        "spectral_distance_db": spec.get("value"),
        "pass": verification.get("pass", False),
        "failed_checks": _failed_checks(verification),
        "output": out_path,
    }

    json_path = os.path.join(out_dir, f"{Path(path).stem}_gate.json")
    with open(json_path, "w") as fh:
        json.dump({"report": report, "row": row}, fh, indent=1)

    return row


def print_summary(rows):
    table = Table(title="Phase C Integration Gate", show_header=True, header_style="bold magenta")
    for col in (
        "Track", "Genre", "Phase A", "Wall(s)", "RSS(MB)",
        "LUFS", "TP(dBTP)", "Spec(dB)", "Pass", "Failed",
    ):
        table.add_column(col)
    for r in rows:
        table.add_row(
            Path(r["file"]).name,
            r["genre"],
            str(r["phase_a"]),
            str(r["wall_clock_s"]),
            str(r["peak_rss_mb"]),
            f"{r['output_lufs']:.1f}" if r.get("output_lufs") is not None else "—",
            f"{r['true_peak_dbtp']:.2f}" if r.get("true_peak_dbtp") is not None else "—",
            f"{r['spectral_distance_db']:.2f}" if r.get("spectral_distance_db") is not None else "—",
            "[green]PASS[/green]" if r["pass"] else "[red]FAIL[/red]",
            r["failed_checks"] or "—",
        )
    console.print(table)
    passed = sum(1 for r in rows if r["pass"])
    console.print(
        f"\n[bold]Summary:[/bold] {passed}/{len(rows)} passed | "
        f"results → {RESULTS_CSV}"
    )


def write_results_csv(rows):
    os.makedirs(RESULTS_DIR, exist_ok=True)
    fields = [
        "file", "genre", "phase_a", "elapsed_s", "wall_clock_s", "peak_rss_mb",
        "output_lufs", "true_peak_dbtp", "spectral_distance_db", "pass", "failed_checks",
    ]
    with open(RESULTS_CSV, "w", newline="") as fh:
        writer = csv.DictWriter(fh, fieldnames=fields, extrasaction="ignore")
        writer.writeheader()
        writer.writerows(rows)


def evaluate_library(
    root_path,
    limit=None,
    genre_override=None,
    refs_dir=None,
    profile_dir=None,
    corpus_dir=None,
    no_stems=False,
    light=False,
):
    profile_dir = profile_dir or os.path.join(RESULTS_DIR, "profiles")
    os.makedirs(profile_dir, exist_ok=True)

    tracks = discover_source_tracks(root_path, limit=limit) if root_path else []
    active_refs = refs_dir or os.path.join("tmp", "refs")
    if not tracks and corpus_dir:
        console.print(
            f"[yellow]No source tracks in {root_path}; using synthetic corpus.[/yellow]"
        )
        tracks = generate_synthetic_corpus(
            corpus_dir, refs_dir=active_refs, force=True)[: limit or 3]
    genres_needed = set()
    if genre_override:
        genres_needed.add(_genre_key(genre_override))
    for path in tracks:
        genres_needed.add(_genre_key(genre_override or infer_genre(path)))
    built = ensure_profiles(
        refs_dir=active_refs,
        profile_dir=profile_dir,
        genres=sorted(genres_needed),
    )
    if built["real"]:
        console.print(f"[cyan]Built reference profiles:[/cyan] {', '.join(built['real'])}")
    if built["builtin"]:
        console.print(
            f"[dim]Builtin fallback profiles:[/dim] {', '.join(built['builtin'])}"
        )
    if limit:
        tracks = tracks[:limit]

    if not tracks:
        console.print("[red]No tracks to evaluate.[/red]")
        return []

    out_dir = os.path.join(RESULTS_DIR, "masters")
    os.makedirs(out_dir, exist_ok=True)
    console.print(f"[bold]Evaluating {len(tracks)} track(s)...[/bold]")

    rows = []
    for path in tracks:
        genre = genre_override or infer_genre(path)
        ensure_profile_for_genre(genre, profile_dir=profile_dir, refs_dir=active_refs)

        console.print(f"\n[white]→ {Path(path).name}[/white] [dim]({genre})[/dim]")
        try:
            row = evaluate_track(
                path, genre, profile_dir, out_dir, refs_dir=active_refs,
                no_stems=no_stems, light=light,
            )
            rows.append(row)
            status = "PASS" if row["pass"] else "FAIL"
            console.print(
                f"  {status} | {row['wall_clock_s']}s wall | "
                f"{row['peak_rss_mb']} MB peak RSS"
            )
        except Exception as exc:
            console.print(f"  [red]Error:[/red] {exc}")
            rows.append({
                "file": path,
                "genre": genre,
                "phase_a": "error",
                "elapsed_s": None,
                "wall_clock_s": None,
                "peak_rss_mb": _peak_rss_mb(),
                "output_lufs": None,
                "true_peak_dbtp": None,
                "spectral_distance_db": None,
                "pass": False,
                "failed_checks": str(exc),
            })

    write_results_csv(rows)
    print_summary(rows)
    return rows


def main():
    parser = argparse.ArgumentParser(description="Phase C integration gate for library tracks")
    parser.add_argument(
        "--library",
        default=DEFAULT_LIBRARY,
        help="Root directory of Suno source WAVs",
    )
    parser.add_argument("--limit", type=int, default=3, help="Max tracks to process")
    parser.add_argument("--genre", default=None, help="Force genre for all tracks")
    parser.add_argument(
        "--refs",
        default=None,
        help="refs/ root (refs/<genre>/*.wav); builds profiles before run",
    )
    parser.add_argument(
        "--corpus",
        default=os.path.join("tmp", "corpus"),
        help="Synthetic corpus directory when library is empty",
    )
    parser.add_argument(
        "--regen-corpus",
        action="store_true",
        help="Regenerate synthetic corpus WAVs before evaluation",
    )
    parser.add_argument(
        "--profile-dir",
        default=None,
        help="Override reference profile output directory",
    )
    parser.add_argument(
        "--no-stems",
        action="store_true",
        help="Skip Phase A stem separation (fast smoke / no demucs)",
    )
    parser.add_argument(
        "--light",
        action="store_true",
        help="Light per-stem processing when stems are enabled",
    )
    args = parser.parse_args()

    refs = args.refs
    if refs is None and os.path.isdir("refs"):
        refs = "refs"

    if args.regen_corpus:
        synth_refs = refs or os.path.join("tmp", "refs")
        generate_synthetic_corpus(args.corpus, refs_dir=synth_refs, force=True)

    evaluate_library(
        args.library if os.path.isdir(args.library) else None,
        limit=args.limit,
        genre_override=args.genre,
        refs_dir=refs,
        profile_dir=args.profile_dir,
        corpus_dir=args.corpus,
        no_stems=args.no_stems,
        light=args.light,
    )


if __name__ == "__main__":
    main()
