#!/usr/bin/python3
import numpy as np
import matplotlib.pyplot as plt
from scipy.io import wavfile
from scipy.signal import butter, sosfiltfilt, peak_widths, find_peaks

def filter_morse_tone(input_wav, output_wav="filtered_morse.wav", bandwidth_hz=100):
    # 1. WAV-Datei einlesen
    sample_rate, data = wavfile.read(input_wav)
    
    # Auf Mono reduzieren (falls Stereo)
    if data.ndim > 1:
        data = data[:, 0]
        
    # In Float konvertieren für präzise Filterung
    data_float = data.astype(np.float64)
    n_samples = len(data_float)
    
    # 2. Dominante Frequenz (Morse-Ton) per FFT finden
    fft_spectrum = np.fft.rfft(data_float)
    frequencies = np.fft.rfftfreq(n_samples, d=1/sample_rate)
    amplitudes = np.abs(fft_spectrum)
    
    # Frequenzbereich einschränken (z. B. typische CW-Töne zwischen 300 Hz und 1500 Hz)
    cw_mask = (frequencies >= 300) & (frequencies <= 1500)
    cw_freqs = frequencies[cw_mask]
    cw_amps = amplitudes[cw_mask]
    
    # Peak mit höchster Amplitude im Bereich finden
    peak_idx = np.argmax(cw_amps)
    cw_center_freq = cw_freqs[peak_idx]
    print(f"[+] Erkannte Morse-Tonfrequenz: {cw_center_freq:.1f} Hz")
    
    # 3. Bandpass-Filter entwerfen (Butterworth SOS für Stabilität)
    lowcut = max(20.0, cw_center_freq - (bandwidth_hz / 2))
    highcut = min(sample_rate / 2 - 1.0, cw_center_freq + (bandwidth_hz / 2))
    
    # 4. Ordnung Butterworth Bandpass
    sos = butter(N=4, Wn=[lowcut, highcut], btype='bandpass', fs=sample_rate, output='sos')
    
    # Zero-Phase-Filterung anwenden (keine Phasenverschiebung im Signal)
    filtered_data = sosfiltfilt(sos, data_float)
    
    # 4. Normalisieren & Speichern
    # Verhindert Clipping beim Speichern im ursprünglichen Datentyp
    max_val = np.max(np.abs(filtered_data))
    if max_val > 0:
        if np.issubdtype(data.dtype, np.integer):
            i_info = np.iinfo(data.dtype)
            filtered_scaled = (filtered_data / max_val * i_info.max * 0.9).astype(data.dtype)
        else:
            filtered_scaled = (filtered_data / max_val * 0.9).astype(data.dtype)
    else:
        filtered_scaled = filtered_data.astype(data.dtype)
        
    wavfile.write(output_wav, sample_rate, filtered_scaled)
    print(f"[+] Gefilterte Datei gespeichert als: {output_wav}")
    
    # 5. Visualisierung (Vorher vs. Nachher)
    filtered_fft = np.abs(np.fft.rfft(filtered_data))
    
    fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(10, 6))
    fig.subplots_adjust(hspace=0.4)
    
    # Spektrum Vorher
    ax1.plot(frequencies, amplitudes / n_samples, color='gray', alpha=0.7, label='Original')
    ax1.axvline(cw_center_freq, color='red', linestyle='--', label=f'CW Tone ({cw_center_freq:.0f} Hz)')
    ax1.set_title("Frequenzspektrum (Original)")
    ax1.set_xlabel("Frequenz [Hz]")
    ax1.set_ylabel("Amplitude")
    ax1.set_xlim(0, 3000)
    ax1.grid(True)
    ax1.legend()
    
    # Spektrum Nachher
    ax2.plot(frequencies, filtered_fft / n_samples, color='green', label='Gefiltert')
    ax2.set_title(f"Frequenzspektrum nach Bandpass ({bandwidth_hz} Hz Bandbreite)")
    ax2.set_xlabel("Frequenz [Hz]")
    ax2.set_ylabel("Amplitude")
    ax2.set_xlim(0, 3000)
    ax2.grid(True)
    ax2.legend()
    
    plt.show()

# Beispielaufruf:
filter_morse_tone("input.wav", bandwidth_hz=10)
