# tools/automaster_app/pipeline/match_eq.py
"""Linear-phase reference match-EQ.

Correction = ref_curve - src_curve, Gaussian-smoothed in octave domain,
capped +-4 dB, zeroed below 30 Hz, cut-only above the source's pre-extension
roll-off (hf_extension owns that region). Applied as 4097-tap windowed-sinc
FIR via fftconvolve, latency compensated.
"""
import numpy as np
import pyloudnorm as pyln
from scipy.signal import firwin2, fftconvolve
from scipy.ndimage import gaussian_filter1d
from .reference_engine import THIRD_OCTAVE_HZ, third_octave_spectrum_db

CAP_DB = 4.0
NTAPS = 4097
LOW_ZERO_HZ = 30.0
VERIFY_LUFS = -14.0
SPECTRAL_SEL = (np.array(THIRD_OCTAVE_HZ) >= 50) & (np.array(THIRD_OCTAVE_HZ) <= 16000)


def _lufs_normalized(audio, sr, target_lufs=VERIFY_LUFS):
    meter = pyln.Meter(sr)
    try:
        lufs = meter.integrated_loudness(audio)
        if np.isfinite(lufs):
            return audio * 10 ** ((target_lufs - lufs) / 20.0)
    except Exception:
        pass
    return audio


def spectral_deviation_db(audio, sr, profile):
    """Per-band deviation vs profile using verify-compatible LUFS-normalized measurement."""
    hz = np.array(profile["third_octave_hz"])
    ref = np.array(profile["third_octave_db"])
    spec = np.array(third_octave_spectrum_db(_lufs_normalized(audio, sr), sr))
    dev = spec - ref
    return hz, dev, float(np.max(np.abs(dev[SPECTRAL_SEL])))


def correction_curve_db(audio, sr, profile, pre_extension_rolloff_hz=None,
                        gaussian_sigma=1.5):
    hz = np.array(profile["third_octave_hz"])
    ref = np.array(profile["third_octave_db"])
    if len(hz) != len(THIRD_OCTAVE_HZ):
        raise ValueError(
            f"profile third_octave_hz length {len(hz)} != {len(THIRD_OCTAVE_HZ)}")
    if len(ref) != len(THIRD_OCTAVE_HZ):
        raise ValueError(
            f"profile third_octave_db length {len(ref)} != {len(THIRD_OCTAVE_HZ)}")
    src = np.array(third_octave_spectrum_db(_lufs_normalized(audio, sr), sr))
    corr = ref - src
    if gaussian_sigma > 0:
        corr = gaussian_filter1d(corr, sigma=gaussian_sigma)
    corr = np.clip(corr, -CAP_DB, CAP_DB)
    corr[hz < LOW_ZERO_HZ] = 0.0
    if pre_extension_rolloff_hz is not None:
        corr[hz > pre_extension_rolloff_hz] = np.minimum(
            corr[hz > pre_extension_rolloff_hz], 0.0)
    return hz, corr


def apply_match_eq(audio, sr, profile, pre_extension_rolloff_hz=None,
                   gaussian_sigma=1.5):
    mono_1d = audio.ndim == 1
    if mono_1d:
        audio = audio[:, None]
    hz, corr = correction_curve_db(
        audio, sr, profile, pre_extension_rolloff_hz, gaussian_sigma=gaussian_sigma)
    nyq = sr / 2.0
    freqs = np.concatenate(([0.0], hz[hz < nyq], [nyq]))
    gains_db = np.concatenate(([corr[0]], corr[hz < nyq], [corr[hz < nyq][-1]]))
    gains = 10 ** (gains_db / 20.0)
    fir = firwin2(NTAPS, freqs, gains, fs=sr)
    delay = (NTAPS - 1) // 2
    padded = np.pad(audio, ((0, delay), (0, 0)))
    out = fftconvolve(padded, fir[:, None], mode="full", axes=0)
    out = out[delay: delay + audio.shape[0]]
    return out[:, 0] if mono_1d else out
