import numpy as np
import matplotlib.pyplot as plt
from scipy import signal
from scipy.fft import fft, fftfreq

# ----------------------------------------------------------------------
# 1. Convolutional Fourier Series
# ----------------------------------------------------------------------
def convolutional_fourier_series(frequencies, amplitudes, duration, sample_rate, phases=None):
    """
    Compute the repeated convolution of cosine waves.
    
    Parameters:
        frequencies : list of float
            Frequencies (Hz) for each cosine component.
        amplitudes : list of float
            Amplitudes for each component.
        duration : float
            Total time (seconds).
        sample_rate : int
            Samples per second.
        phases : list of float, optional
            Phase shifts (radians) for each component. Defaults to zero.
    
    Returns:
        t : np.ndarray
            Time axis.
        result : np.ndarray
            The convolved signal (Memory Echo Resonance).
    """
    if phases is None:
        phases = [0.0] * len(frequencies)
    
    # Create time axis
    t = np.linspace(0, duration, int(duration * sample_rate), endpoint=False)
    
    # Build the list of cosine signals
    signals = []
    for f, A, phi in zip(frequencies, amplitudes, phases):
        sig = A * np.cos(2 * np.pi * f * t + phi)
        signals.append(sig)
    
    # Start with the first signal
    result = signals[0]
    # Repeatedly convolve with the remaining signals
    for i in range(1, len(signals)):
        # Convolution preserves length? We'll keep the valid range
        # Use 'same' mode to keep length consistent with original (approximately)
        result = signal.convolve(result, signals[i], mode='same')
    
    return t, result

# ----------------------------------------------------------------------
# 2. Helper: plot time domain and spectrum
# ----------------------------------------------------------------------
def plot_signal_and_spectrum(t, y, title):
    fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(10, 6))
    
    # Time domain
    ax1.plot(t, y)
    ax1.set_xlabel('Time (s)')
    ax1.set_ylabel('Amplitude')
    ax1.set_title(f'{title} - Time Domain')
    ax1.grid(True)
    
    # Frequency domain (magnitude spectrum)
    n = len(y)
    yf = fft(y)
    xf = fftfreq(n, t[1] - t[0])[:n//2]
    magnitude = np.abs(yf)[:n//2]
    ax2.plot(xf, magnitude)
    ax2.set_xlabel('Frequency (Hz)')
    ax2.set_ylabel('Magnitude')
    ax2.set_title(f'{title} - Frequency Spectrum')
    #ax2.set_xlim(0, max(10, max(xf)*2))  # adjust view
    ax2.grid(True)
    
    plt.tight_layout()
    plt.show()

# ----------------------------------------------------------------------
# 3. Examples
# ----------------------------------------------------------------------
if __name__ == "__main__":
    sample_rate = 1000  # Hz
    duration = 2.0      # seconds
    
    # Example 1: Two frequencies (beats via convolution)
    print("Example 1: Two frequencies – memory echo resonance")
    freqs1 = [5, 8]
    amps1 = [1.0, 1.0]
    t1, y1 = convolutional_fourier_series(freqs1, amps1, duration, sample_rate)
    plot_signal_and_spectrum(t1, y1, "Convolutional Fourier: 5Hz ⊗ 8Hz")
    
    # Example 2: Three frequencies – harmonic series
    print("Example 2: Harmonic series 1,2,3 Hz")
    freqs2 = [1, 2, 3]
    amps2 = [1.0, 0.8, 0.6]
    t2, y2 = convolutional_fourier_series(freqs2, amps2, duration, sample_rate)
    plot_signal_and_spectrum(t2, y2, "Convolutional Fourier: 1⊗2⊗3 Hz")
    
    # Example 3: Inharmonic set – chaotic echo
    print("Example 3: Inharmonic set 1, √2, π Hz")
    freqs3 = [1, np.sqrt(2), np.pi]
    amps3 = [1.0, 1.0, 1.0]
    t3, y3 = convolutional_fourier_series(freqs3, amps3, duration, sample_rate)
    plot_signal_and_spectrum(t3, y3, "Convolutional Fourier: 1 ⊗ √2 ⊗ π Hz")
    
    # Example 4: With different phases
    print("Example 4: Two frequencies with phase shifts (90°, 180°)")
    freqs4 = [2, 5]
    amps4 = [1.0, 1.0]
    phases4 = [np.pi/2, np.pi]
    t4, y4 = convolutional_fourier_series(freqs4, amps4, duration, sample_rate, phases4)
    plot_signal_and_spectrum(t4, y4, "Convolutional Fourier: 2Hz(90°) ⊗ 5Hz(180°)")
