
import unittest
import numpy as np
from engine.processor import AudioEngine

class TestStereo(unittest.TestCase):
    def setUp(self):
        self.engine = AudioEngine()
        self.engine.sample_rate = 44100

    def test_ms_conversion_perfect_reconstruction(self):
        """Test M/S encode/decode round trip."""
        sr = 44100
        # Create random stereo signal
        np.random.seed(42)
        original = np.random.normal(0, 0.5, (sr, 2))
        
        # Manually do M/S to test logic
        mid = (original[:, 0] + original[:, 1]) / 2
        side = (original[:, 0] - original[:, 1]) / 2
        
        left_rec = mid + side
        right_rec = mid - side
        
        reconstructed = np.column_stack((left_rec, right_rec))
        
        # Check equality
        np.testing.assert_allclose(reconstructed, original, atol=1e-7)

    def test_stereo_width_control(self):
        """Test apply_stereo_imaging width control."""
        sr = 44100
        t = np.linspace(0, 1.0, sr, endpoint=False)
        
        # Create a signal completely within Mid band (1kHz)
        # Mid content: 1kHz sine
        mid_sig = np.sin(2 * np.pi * 1000 * t)
        # Side content: 1kHz sine (phase shifted or just uncorrelated sine)
        # Use a cosine for side to be orthogonal/different
        side_sig = np.cos(2 * np.pi * 1000 * t) * 0.5
        
        left = mid_sig + side_sig
        right = mid_sig - side_sig
        stereo = np.column_stack((left, right))
        
        # 1. Test Mono (Width = 0)
        mono_processed = self.engine.apply_stereo_imaging(
            stereo,
            low_width=0.0, mid_width=0.0, high_width=0.0
        )
        
        # Left should equal Right
        np.testing.assert_allclose(mono_processed[:, 0], mono_processed[:, 1], atol=1e-5)
        
        # 2. Test Widening (Width = 2.0)
        # Apply to Mid band
        wide_processed = self.engine.apply_stereo_imaging(
            stereo,
            low_width=1.0, mid_width=2.0, high_width=1.0
        )
        
        # Calculate new side energy
        new_mid = (wide_processed[:, 0] + wide_processed[:, 1]) / 2
        new_side = (wide_processed[:, 0] - wide_processed[:, 1]) / 2
        
        rms_orig_side = np.sqrt(np.mean(side_sig**2))
        rms_new_side = np.sqrt(np.mean(new_side**2))
        
        print(f"Original Side RMS: {rms_orig_side:.4f}")
        print(f"New Side RMS: {rms_new_side:.4f}")
        
        # Should be almost double (leakage might reduce it slightly)
        self.assertAlmostEqual(rms_new_side, rms_orig_side * 2, delta=0.05)

    def test_frequency_dependent_width(self):
        """Test different widths for different frequencies."""
        sr = 44100
        t = np.linspace(0, 1.0, sr, endpoint=False)
        
        # Low frequency (100Hz)
        low = np.sin(2 * np.pi * 100 * t)
        # High frequency (10kHz)
        high = np.sin(2 * np.pi * 10000 * t)
        
        # Make them stereo (out of phase side content)
        side_low = low * 0.5
        side_high = high * 0.5
        
        left = (low + side_low) + (high + side_high)
        right = (low - side_low) + (high - side_high)
        stereo = np.column_stack((left, right))
        
        # Process: Low -> Mono (0.0), High -> Wide (2.0)
        processed = self.engine.apply_stereo_imaging(
            stereo,
            low_width=0.0,
            mid_width=1.0,
            high_width=2.0
        )
        
        # Analyze Low (100Hz) - Should be mono
        # We can just check that at 100Hz, L ~= R
        # FFT analysis
        def get_lr_magnitude_diff_at_freq(sig, freq):
            l_fft = np.fft.rfft(sig[:, 0])
            r_fft = np.fft.rfft(sig[:, 1])
            freqs = np.fft.rfftfreq(len(sig), 1/sr)
            idx = np.argmin(np.abs(freqs - freq))
            
            # Phase difference or just side energy?
            # If mono, L and R are identical, so L-R (Side) is 0
            side_fft = (l_fft - r_fft) / 2
            return np.abs(side_fft[idx])
            
        side_low_mag = get_lr_magnitude_diff_at_freq(processed, 100)
        side_high_mag = get_lr_magnitude_diff_at_freq(processed, 10000)
        
        print(f"Side Mag @ 100Hz (Target 0): {side_low_mag:.4f}")
        print(f"Side Mag @ 10kHz (Target High): {side_high_mag:.4f}")
        
        # Check relative reduction
        # Original 100Hz side magnitude should be high (~11000)
        # Reduced should be < 10% of that
        self.assertLess(side_low_mag, 2000.0)
        self.assertGreater(side_high_mag, 10000.0)

if __name__ == '__main__':
    unittest.main()
