#!/usr/bin/env catnip
# Débruitage audio et spectrogrammes avec SciPy.
#
# Un signal harmonique modulé simule une voix simple. On lui ajoute un ronflement
# secteur à 50 Hz et un bruit large bande reproductible. Catnip diffuse trois
# presets de filtres sur le même signal ; scipy.signal effectue le filtrage sans
# phase et l'analyse fréquentielle. Le meilleur résultat est écrit en WAV et une
# planche compare les spectrogrammes avant/après.
#
# Contrôles : amélioration du SNR contre le signal propre connu, atténuation du
# bin 50 Hz et roundtrip du WAV filtré après relecture depuis le disque.
#
# DEPS: numpy scipy pillow

numpy = import('numpy')
signal = import('scipy.signal')
wavfile = import('scipy.io.wavfile')
Image = import('PIL.Image')
ImageDraw = import('PIL.ImageDraw')
math = import('math')
import('pathlib', 'Path')

script_dir = Path(META.file).parent
output_dir = script_dir / 'output'
output_dir.mkdir(exist_ok=True)

SAMPLE_RATE = 16000
DURATION = 3.0
N_SAMPLES = int(SAMPLE_RATE * DURATION)

struct FilterPreset {
    name: str; low_hz: float; high_hz: float; notch_hz: float;
    quality: float; order: int;
}

struct AudioResult {
    preset: FilterPreset; waveform; snr_db: float; rms_error: float;

    display(self) => {
        notch = if self.preset.notch_hz > 0.0 { f", notch={self.preset.notch_hz} Hz" } else { "" }
        print(
            f"  {self.preset.name:<16} : SNR={round(self.snr_db, 2):>6} dB  " +
                f"RMSE={round(self.rms_error, 5)}{notch}"
        )
    }
}

time = numpy.arange(N_SAMPLES, dtype='float64') / SAMPLE_RATE
envelope = 0.65 + 0.35 * numpy.square(numpy.sin(2.0 * math.pi * 2.0 * time))
clean = envelope *
    (0.45 * numpy.sin(2.0 * math.pi * 220.0 * time) + 0.22 * numpy.sin(2.0 * math.pi * 440.0 * time) +
        0.12 * numpy.sin(2.0 * math.pi * 880.0 * time))

rng = numpy.random.default_rng(42)
hum = 0.18 * numpy.sin(2.0 * math.pi * 50.0 * time)
wideband_noise = rng.normal(0.0, 0.12, N_SAMPLES)
noisy = clean + hum + wideband_noise

snr_db = (reference, candidate): float => {
    signal_power = float(numpy.sum(numpy.square(reference)))
    noise_power = float(numpy.sum(numpy.square(candidate - reference)))
    10.0 * math.log10(signal_power / noise_power)
}

presets = list(
    FilterPreset('large bande', 30.0, 1400.0, 0.0, 30.0, 4),
    FilterPreset('notch 50 Hz', 30.0, 1400.0, 50.0, 30.0, 4),
    FilterPreset('voix étroite', 180.0, 1000.0, 0.0, 30.0, 5),
)

apply_filter = (preset: FilterPreset): AudioResult => {
    filtered = noisy
    if preset.notch_hz > 0.0 {
        coeffs = signal.iirnotch(preset.notch_hz, preset.quality, fs=SAMPLE_RATE)
        filtered = signal.filtfilt(coeffs[0], coeffs[1], filtered)
    }

    sos = signal.butter(
        preset.order,
        list(preset.low_hz, preset.high_hz),
        btype='bandpass',
        fs=SAMPLE_RATE,
        output='sos',
    )
    filtered = signal.sosfiltfilt(sos, filtered)
    error = filtered - clean
    AudioResult(
        preset,
        filtered,
        snr_db(clean, filtered),
        float(numpy.sqrt(numpy.mean(numpy.square(error)))),
    )
}

noisy_snr = snr_db(clean, noisy)
print("⇒ Signal synthétique reproductible")
print(f"  échantillonnage={SAMPLE_RATE} Hz  durée={DURATION} s  échantillons={N_SAMPLES}")
print(f"  SNR bruité={round(noisy_snr, 2)} dB")

print()
print("⇒ Diffusion de trois presets de filtrage")
results = presets.[(preset) => { apply_filter(preset) }]
for result in results {
    result.display()
}

best = max(results, key=(result) => { result.snr_db })
improvement = best.snr_db - noisy_snr

# Contrôle spectral indépendant : amplitude du bin FFT le plus proche de 50 Hz.
freqs = numpy.fft.rfftfreq(N_SAMPLES, 1.0 / SAMPLE_RATE)
hum_idx = int(numpy.argmin(numpy.abs(freqs - 50.0)))
noisy_spectrum = numpy.abs(numpy.fft.rfft(noisy))
best_spectrum = numpy.abs(numpy.fft.rfft(best.waveform))
hum_attenuation = 20.0 *
    math.log10(
        (float(noisy_spectrum[hum_idx]) + 1.0e-12) / (float(best_spectrum[hum_idx]) + 1.0e-12)
    )

print()
print(f"⇒ Meilleur preset : {best.preset.name}")
print(f"  gain de SNR={round(improvement, 2)} dB")
print(f"  atténuation mesurée à 50 Hz={round(hum_attenuation, 2)} dB")

# WAV float32 : la version bruitée est bornée à [-1, 1] pour l'écoute ; le
# résultat filtré reste déjà dans cette plage, mais on applique le même contrat.
noisy_pcm = numpy.asarray(numpy.clip(noisy, -1.0, 1.0), dtype='float32')
best_pcm = numpy.asarray(numpy.clip(best.waveform, -1.0, 1.0), dtype='float32')
noisy_path = output_dir / 'audio_noisy.wav'
filtered_path = output_dir / 'audio_denoised.wav'
wavfile.write(str(noisy_path), SAMPLE_RATE, noisy_pcm)
wavfile.write(str(filtered_path), SAMPLE_RATE, best_pcm)

# Roundtrip disque : fréquence, dimensions et valeurs doivent être conservées.
reloaded = wavfile.read(str(filtered_path))
roundtrip_rate = int(reloaded[0])
roundtrip_audio = numpy.asarray(reloaded[1], dtype='float32')
roundtrip_error = float(numpy.max(numpy.abs(roundtrip_audio - best_pcm)))
roundtrip_ok = roundtrip_rate == SAMPLE_RATE and roundtrip_audio.shape == best_pcm.shape and roundtrip_error == 0.0

print()
print(f"⇒ WAV bruité → {noisy_path}")
print(f"⇒ WAV filtré → {filtered_path}")
print(f"  roundtrip WAV exact={roundtrip_ok}, erreur max={roundtrip_error}")

# Spectrogrammes calculés avec une même fenêtre et une même échelle en dB. La
# palette est construite avec numpy puis convertie en image Pillow.
spectrogram = (waveform) => {
    spec = signal.spectrogram(
        waveform,
        fs=SAMPLE_RATE,
        window='hann',
        nperseg=512,
        noverlap=384,
        scaling='spectrum',
        mode='magnitude',
    )
    limit = int(numpy.searchsorted(spec[0], 2000.0))
    tuple(spec[0][:limit], spec[1], spec[2][:limit])
}

noisy_spec = spectrogram(noisy)
best_spec = spectrogram(best.waveform)
reference_db = float(numpy.max(20.0 * numpy.log10(numpy.maximum(noisy_spec[2], 1.0e-10))))

spec_image = (magnitude) => {
    db = 20.0 * numpy.log10(numpy.maximum(magnitude, 1.0e-10))
    level = numpy.clip((db - (reference_db - 65.0)) / 65.0, 0.0, 1.0)
    red = numpy.clip(255.0 * (1.5 * level - 0.25), 0.0, 255.0)
    green = numpy.clip(255.0 * numpy.sqrt(level), 0.0, 255.0)
    blue = numpy.clip(210.0 * (1.0 - level) + 40.0 * level, 0.0, 255.0)
    rgb = numpy.stack(list(red, green, blue), axis=2).astype('uint8')
    Image.fromarray(numpy.flipud(rgb), 'RGB').resize(tuple(600, 270), Image.Resampling.BILINEAR)
}

left = spec_image(noisy_spec[2])
right = spec_image(best_spec[2])
board = Image.new('RGB', tuple(1240, 350), tuple(18, 18, 24))
board.paste(left, tuple(12, 55))
board.paste(right, tuple(628, 55))
draw = ImageDraw.Draw(board)
draw.text(tuple(12, 14), f"Bruit + ronflement — SNR {round(noisy_snr, 2)} dB", fill=tuple(235, 235, 240))
draw.text(
    tuple(628, 14),
    f"{best.preset.name} — SNR {round(best.snr_db, 2)} dB",
    fill=tuple(235, 235, 240),
)
draw.text(tuple(12, 38), "2 kHz", fill=tuple(190, 190, 200))
draw.text(tuple(12, 330), "0 Hz", fill=tuple(190, 190, 200))
draw.text(tuple(628, 38), "2 kHz", fill=tuple(190, 190, 200))
draw.text(tuple(628, 330), "0 Hz", fill=tuple(190, 190, 200))

image_path = output_dir / 'audio_denoising_spectrogram.png'
board.save(str(image_path))

oracle_ok = improvement > 10.0 and hum_attenuation > 20.0 and roundtrip_ok
print(f"⇒ Oracle global (SNR + 50 Hz + roundtrip) : {oracle_ok}")
print(f"⇒ Spectrogrammes → {image_path}")

# Aperçu navigateur : sert l'image une fois puis rend la main.
# --no-browser garde un chemin headless (le PNG reste écrit ci-dessus).
if '--no-browser' not in import('sys').argv {
    http = import('http')
    b64 = import('base64').b64encode(image_path.read_bytes()).decode('ascii')
    http.serve(f'<!doctype html><meta charset="utf-8"><title>Audio denoising spectrogram</title><body style="margin:0;background:#0d1117"><img style="max-width:100%;display:block;margin:0 auto" src="data:image/png;base64,{b64}">',
        0, 'text/html; charset=utf-8', True)
}