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

# ----------------------------------------------------------------------
# 1. Geometric Mean Fourier Series
# ----------------------------------------------------------------------
def geometric_mean_fourier_series(frequencies, amplitudes, duration, sample_rate, phases=None, eps=1e-12):
    """
    Compute the geometric mean across sinusoids at each time point.
    
    f(t) = sign_product * (∏ |A_k cos(2π f_k t + φ_k)|)^(1/N)
    
    Parameters:
        frequencies : list of float
            Frequencies (Hz) for each 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.
        eps : float
            Small constant to avoid log(0).
    
    Returns:
        t : np.ndarray
            Time axis.
        result : np.ndarray
            The geometric‑mean waveform (log‑compressed resonance).
    """
    if phases is None:
        phases = [0.0] * len(frequencies)
    
    # Time axis
    t = np.linspace(0, duration, int(duration * sample_rate), endpoint=False)
    
    # Build a matrix: rows = time, columns = signal index
    signals = np.zeros((len(t), len(frequencies)))
    for i, (f, A, phi) in enumerate(zip(frequencies, amplitudes, phases)):
        signals[:, i] = A * np.cos(2 * np.pi * f * t + phi)
    
    # Geometric mean across columns (the signal dimension)
    # Compute product of absolute values, then 1/N power
    abs_prod = np.prod(np.abs(signals) + eps, axis=1)
    sign_prod = np.sign(np.prod(signals, axis=1))
    geom_mean = sign_prod * (abs_prod ** (1.0 / len(frequencies)))
    
    return t, geom_mean

# ----------------------------------------------------------------------
# 2. Helper: plot time domain and spectrum
# ----------------------------------------------------------------------
def plot_signal_and_spectrum(t, y, title, max_freq=None):
    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')
    if max_freq is None:
        max_freq = max(10, max(frequencies)*2) if 'frequencies' in dir() else 10
    ax2.set_xlim(0, max_freq)
    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 – balanced compression
    print("Example 1: Geometric mean of 5Hz and 8Hz")
    freqs1 = [5, 8]
    amps1 = [1.0, 1.0]
    t1, y1 = geometric_mean_fourier_series(freqs1, amps1, duration, sample_rate)
    plot_signal_and_spectrum(t1, y1, "Geometric Mean: 5Hz & 8Hz", max_freq=20)
    
    # Example 2: Three frequencies – harmonic series with different amplitudes
    print("Example 2: Harmonic series 1,2,3 Hz with amplitudes 1.0, 0.5, 0.2")
    freqs2 = [1, 2, 3]
    amps2 = [1.0, 0.5, 0.2]
    t2, y2 = geometric_mean_fourier_series(freqs2, amps2, duration, sample_rate)
    plot_signal_and_spectrum(t2, y2, "Geometric Mean: 1,2,3 Hz (unequal amplitudes)", max_freq=10)
    
    # Example 3: With phase shifts – effect on waveform shape (spectrum magnitude unchanged)
    print("Example 3: 2Hz and 5Hz with phases 90° and 180°")
    freqs3 = [2, 5]
    amps3 = [1.0, 1.0]
    phases3 = [np.pi/2, np.pi]
    t3, y3 = geometric_mean_fourier_series(freqs3, amps3, duration, sample_rate, phases3)
    plot_signal_and_spectrum(t3, y3, "Geometric Mean: 2Hz(90°) & 5Hz(180°)", max_freq=15)
    
    # Example 4: Inharmonic set – log‑compressed chaos
    print("Example 4: Inharmonic set 1, √2, π Hz")
    freqs4 = [1, np.sqrt(2), np.pi]
    amps4 = [1.0, 1.0, 1.0]
    t4, y4 = geometric_mean_fourier_series(freqs4, amps4, duration, sample_rate)
    plot_signal_and_spectrum(t4, y4, "Geometric Mean: 1, √2, π Hz", max_freq=8)
    
    # Example 5: Very different amplitudes – outlier suppression
    print("Example 5: Strong 10Hz (A=2) and weak 3Hz (A=0.1)")
    freqs5 = [10, 3]
    amps5 = [2.0, 0.1]
    t5, y5 = geometric_mean_fourier_series(freqs5, amps5, duration, sample_rate)
    plot_signal_and_spectrum(t5, y5, "Geometric Mean: strong 10Hz + weak 3Hz", max_freq=15)