#!/usr/bin/env python3
"""
freqresp.py — Measure audio loopback frequency response (20 Hz – 20 kHz)

Usage:
    python3 freqresp.py [device]        # default: hw:0,0
    python3 freqresp.py hw:0,0

Plays stepped sine tones through the output, records from the input,
computes amplitude via FFT, prints a text table, and saves freqresp.png.
"""

import numpy as np
import sys
import os
import subprocess

DEVICE  = sys.argv[1] if len(sys.argv) > 1 else "hw:0,0"
RATE    = 48000
CH      = 2
FMT     = "S32_LE"
BPS     = 4          # bytes per sample per channel (int32)
TONE_T  = 1.0        # tone duration seconds
SETTLE  = 0.2        # discard first N seconds of each tone (settling)
GAP_T   = 0.2        # silence before first tone (record sync)
FREQS   = np.logspace(np.log10(20), np.log10(20000), 40).tolist()


def get_audio_blocking_services():
    """Detect running audio services and processes that may block ALSA device access."""
    processes = []

    # Try using lsof to find processes holding ALSA devices
    try:
        result = subprocess.run(["lsof", "+D", "/dev/snd"],
                               capture_output=True, text=True, timeout=5)
        if result.returncode == 0 and result.stdout:
            for line in result.stdout.split('\n')[1:]:  # skip header
                if not line.strip():
                    continue
                parts = line.split()
                if len(parts) >= 2:
                    cmd = parts[0]
                    pid = parts[1]
                    seen = False
                    # Map common process names to friendly names
                    if cmd in ("pipewire", "pipewire-pulse"):
                        if not any(p[0] == "PipeWire" for p in processes):
                            processes.append(("PipeWire", pid, cmd))
                            seen = True
                    elif cmd == "pulseaudio":
                        if not any(p[0] == "PulseAudio" for p in processes):
                            processes.append(("PulseAudio", pid, cmd))
                            seen = True
                    elif cmd == "jackd":
                        if not any(p[0] == "JACK" for p in processes):
                            processes.append(("JACK", pid, cmd))
                            seen = True
                    elif "alsa" in cmd.lower():
                        if not any(p[2] == cmd for p in processes):
                            processes.append(("ALSA", pid, cmd))
                            seen = True
                    if seen and processes:
                        return processes
    except Exception:
        pass

    # Fallback: check /proc for audio-related processes
    try:
        result = subprocess.run(["pgrep", "-l", "-f", "pipewire|pulseaudio|jackd|alsa"],
                               capture_output=True, text=True, timeout=5)
        if result.returncode == 0 and result.stdout:
            for line in result.stdout.split('\n'):
                if not line.strip():
                    continue
                parts = line.split(maxsplit=1)
                if len(parts) >= 2:
                    pid = parts[0]
                    cmd = parts[1]
                    if "pipewire" in cmd.lower():
                        if not any(p[0] == "PipeWire" for p in processes):
                            processes.append(("PipeWire", pid, cmd.split()[0]))
                    elif "pulseaudio" in cmd.lower():
                        if not any(p[0] == "PulseAudio" for p in processes):
                            processes.append(("PulseAudio", pid, cmd.split()[0]))
                    elif "jackd" in cmd.lower():
                        if not any(p[0] == "JACK" for p in processes):
                            processes.append(("JACK", pid, cmd.split()[0]))
            if processes:
                return processes
    except Exception:
        pass

    # Fallback: check systemctl for running services
    common_services = {
        "pulseaudio": "PulseAudio",
        "pipewire": "PipeWire",
        "jackd": "JACK",
        "alsa-state": "ALSA State",
    }

    try:
        result = subprocess.run(["systemctl", "list-units", "--all", "--plain"],
                               capture_output=True, text=True, timeout=5)
        running = result.stdout
        for cmd, name in common_services.items():
            if cmd in running or f"{cmd}.service" in running:
                processes.append((name, "?", cmd))
    except Exception:
        pass

    return processes


def make_tone(freq, duration, amplitude=0.5):
    n = int(RATE * duration)
    t = np.linspace(0, duration, n, endpoint=False)
    mono = (np.sin(2 * np.pi * freq * t) * amplitude * (2**31 - 1)).astype(np.int32)
    stereo = np.empty(n * CH, dtype=np.int32)
    stereo[0::2] = mono
    stereo[1::2] = mono
    return stereo.tobytes()


def record_and_play(play_signal, total_dur, device):
    import alsaaudio
    import threading

    period = 1024
    period_bytes = period * CH * BPS
    total_record_bytes = int(RATE * total_dur) * CH * BPS

    errors = []

    def _play():
        try:
            out_pcm = alsaaudio.PCM(alsaaudio.PCM_PLAYBACK, alsaaudio.PCM_NORMAL,
                                    device=device, channels=CH, rate=RATE,
                                    format=alsaaudio.PCM_FORMAT_S32_LE,
                                    periodsize=period)
            pos = 0
            while pos < len(play_signal):
                chunk = play_signal[pos:pos + period_bytes]
                out_pcm.write(chunk)
                pos += len(chunk)
            out_pcm.close()
        except Exception as e:
            errors.append(("playback", e))

    def _record():
        try:
            in_pcm = alsaaudio.PCM(alsaaudio.PCM_CAPTURE, alsaaudio.PCM_NORMAL,
                                   device=device, channels=CH, rate=RATE,
                                   format=alsaaudio.PCM_FORMAT_S32_LE,
                                   periodsize=period)
            recorded = bytearray()
            while len(recorded) < total_record_bytes:
                n, data = in_pcm.read()
                if n > 0:
                    recorded.extend(data)
            in_pcm.close()
            result.append(bytes(recorded))
        except Exception as e:
            errors.append(("recording", e))

    result = []
    rec_t = threading.Thread(target=_record)
    play_t = threading.Thread(target=_play)
    rec_t.start()
    play_t.start()
    play_t.join()
    rec_t.join()

    if errors:
        # Re-raise the first error encountered
        raise errors[0][1]

    if not result:
        raise RuntimeError("Recording failed: no data captured")

    return result[0]


def record_and_play_with_progress(freqs, device):
    """Record each frequency separately to show progress."""
    import alsaaudio
    import threading

    period = 1024
    period_bytes = period * CH * BPS
    all_recorded = bytearray()

    # Initial silence
    silence_dur = GAP_T
    silence_samples = int(RATE * silence_dur) * CH
    silence = (np.zeros(silence_samples, dtype=np.int32)).tobytes()
    all_recorded.extend(silence)

    for idx, freq in enumerate(freqs):
        print(f"  [{idx+1:2d}/{len(freqs)}]  {freq:7.1f} Hz", end="", flush=True)

        # Create tone for this frequency
        tone_signal = make_tone(freq, TONE_T)
        total_dur = TONE_T + 0.1  # slight buffer
        total_record_bytes = int(RATE * total_dur) * CH * BPS

        errors = []

        def _play():
            try:
                out_pcm = alsaaudio.PCM(alsaaudio.PCM_PLAYBACK, alsaaudio.PCM_NORMAL,
                                        device=device, channels=CH, rate=RATE,
                                        format=alsaaudio.PCM_FORMAT_S32_LE,
                                        periodsize=period)
                pos = 0
                while pos < len(tone_signal):
                    chunk = tone_signal[pos:pos + period_bytes]
                    out_pcm.write(chunk)
                    pos += len(chunk)
                out_pcm.close()
            except Exception as e:
                errors.append(("playback", e))

        def _record():
            try:
                in_pcm = alsaaudio.PCM(alsaaudio.PCM_CAPTURE, alsaaudio.PCM_NORMAL,
                                       device=device, channels=CH, rate=RATE,
                                       format=alsaaudio.PCM_FORMAT_S32_LE,
                                       periodsize=period)
                recorded = bytearray()
                while len(recorded) < total_record_bytes:
                    n, data = in_pcm.read()
                    if n > 0:
                        recorded.extend(data)
                in_pcm.close()
                result.append(bytes(recorded))
            except Exception as e:
                errors.append(("recording", e))

        result = []
        rec_t = threading.Thread(target=_record)
        play_t = threading.Thread(target=_play)
        rec_t.start()
        play_t.start()
        play_t.join()
        rec_t.join()

        if errors:
            print(" [FAILED]")
            raise errors[0][1]

        if result:
            all_recorded.extend(result[0])
            print(" ✓")
        else:
            print(" [NO DATA]")
            raise RuntimeError(f"Recording failed for {freq:.1f} Hz: no data captured")

    return bytes(all_recorded)


def analyze(raw, freqs):
    samples = np.frombuffer(raw, dtype=np.int32).astype(np.float64) / (2**31)
    left = samples[0::2]  # left channel only

    pos = int(RATE * GAP_T)
    amplitudes = []

    for freq in freqs:
        tone_n   = int(RATE * TONE_T)
        settle_n = int(RATE * SETTLE)
        chunk = left[pos + settle_n : pos + tone_n]
        pos += tone_n

        if len(chunk) < 512:
            amplitudes.append(None)
            continue

        # Snap to an integer number of complete cycles to eliminate leakage
        cycles = max(1, round(len(chunk) * freq / RATE))
        n_samples = int(round(cycles * RATE / freq))
        n_samples = min(n_samples, len(chunk))
        chunk = chunk[:n_samples]

        # Single-frequency DFT at the exact target frequency — no bin-alignment needed
        n = np.arange(n_samples)
        amp = 2.0 * abs(np.dot(chunk, np.exp(-2j * np.pi * freq * n / RATE))) / n_samples
        amplitudes.append(amp)

    # Normalize: find amplitude closest to 1 kHz as 0 dB reference
    ref_idx = np.argmin(np.abs(np.array(freqs) - 1000.0))
    ref_amp = amplitudes[ref_idx]
    if ref_amp is None or ref_amp < 1e-10:
        ref_amp = max((a for a in amplitudes if a is not None), default=1.0)

    results = []
    for freq, amp in zip(freqs, amplitudes):
        if amp is not None and amp > 1e-10:
            db = 20.0 * np.log10(amp / ref_amp)
        else:
            db = None
        results.append((freq, db))

    return results


def print_table(results):
    print(f"\n{'Freq (Hz)':>12}  {'dB re 1kHz':>12}")
    print("-" * 28)
    for freq, db in results:
        if db is not None:
            print(f"{freq:12.1f}  {db:+12.2f}")
        else:
            print(f"{freq:12.1f}  {'--':>12}")


def save_plot(results, path="freqresp.png"):
    try:
        import matplotlib
        matplotlib.use("Agg")
        import matplotlib.pyplot as plt

        valid = [(f, db) for f, db in results if db is not None]
        if not valid:
            print("No valid data to plot.")
            return
        fv, dv = zip(*valid)

        plt.figure(figsize=(11, 5))
        plt.semilogx(fv, dv, "b-o", markersize=4, linewidth=1.2)
        plt.axhline(0, color="gray", linewidth=0.8, linestyle="--")
        plt.xlabel("Frequency (Hz)")
        plt.ylabel("Level (dB, ref 1 kHz)")
        plt.title(f"Frequency Response — {DEVICE}")
        plt.grid(True, which="both", alpha=0.35)
        plt.xlim(20, 20000)
        plt.ylim(min(dv) - 3, max(dv) + 3)
        plt.xticks(
            [20, 50, 100, 200, 500, 1000, 2000, 5000, 10000, 20000],
            ["20", "50", "100", "200", "500", "1k", "2k", "5k", "10k", "20k"],
        )
        plt.tight_layout()
        plt.savefig(path, dpi=150)
        print(f"\nPlot saved to {os.path.abspath(path)}")
    except ImportError:
        print("\n(matplotlib not installed — skipping plot)")


if __name__ == "__main__":
    print(f"Frequency response measurement")
    print(f"  Device : {DEVICE}")
    print(f"  Points : {len(FREQS)}  (20 Hz – 20 kHz, log spaced)")
    print(f"  Tone   : {TONE_T*1000:.0f} ms each  ({SETTLE*1000:.0f} ms settle)")
    print(f"  Total  : ~{GAP_T + len(FREQS) * TONE_T:.1f} s")
    print("Recording...\n")

    try:
        raw = record_and_play_with_progress(FREQS, DEVICE)
        results = analyze(raw, FREQS)

        print_table(results)
        save_plot(results)
    except Exception as e:
        import alsaaudio
        error_str = str(e)
        if "Device or resource busy" in error_str or isinstance(e, alsaaudio.ALSAAudioError):
            print(f"\nError: Cannot access audio device '{DEVICE}'")
            print(f"  Reason: Device is already in use\n")

            blocking_processes = get_audio_blocking_services()
            if blocking_processes:
                print(f"Blocking processes:")
                for service, pid, cmd in blocking_processes:
                    pid_str = f" (PID {pid})" if pid != "?" else ""
                    print(f"  • {service}{pid_str}: {cmd}")
                print(f"\nTo free the device:")
                for service, pid, cmd in blocking_processes:
                    if service == "PipeWire":
                        print(f"  sudo systemctl stop pipewire pipewire-pulse")
                    elif service == "PulseAudio":
                        print(f"  sudo systemctl stop pulseaudio")
                    elif service == "JACK":
                        print(f"  killall jackd")
                    else:
                        pid_str = f" kill {pid}" if pid != "?" else ""
                        print(f"  {cmd}: sudo systemctl stop {cmd}{pid_str}")
            else:
                print(f"Could not detect blocking services.")
                print(f"Try stopping PipeWire/PulseAudio:")
                print(f"  sudo systemctl stop pipewire pipewire-pulse")

            sys.exit(1)

        # Re-raise other exceptions
        raise
