import numpy as np
from scipy.signal import butter, sosfilt
from .base import BaseProcessorModule

class SidechainModule(BaseProcessorModule):
    """Surgical Sidechain Module with Hold and Multiband Support."""
    
    def process(self, bass_array: np.ndarray, kick_array: np.ndarray, 
                strength: float = 0.5, attack_ms: float = 2.0, 
                hold_ms: float = 10.0, release_ms: float = 100.0, 
                mode: str = "broadband") -> np.ndarray:
        
        sr = self.sample_rate
        
        # 1. Generate Ducking Envelope
        if kick_array.ndim > 1: kick_mono = np.mean(kick_array, axis=1)
        else: kick_mono = kick_array
        
        window_size = int(sr * 0.01)
        envelope = self._calculate_rms_envelope(kick_mono, window_size)
        
        # Attack/Hold/Release samples
        att_s = int(sr * attack_ms / 1000)
        hold_s = int(sr * hold_ms / 1000)
        rel_s = int(sr * release_ms / 1000)
        
        smoothed_env = np.zeros_like(envelope)
        g_att = 1.0 - np.exp(-1.0 / (att_s + 1e-9))
        g_rel = 1.0 - np.exp(-1.0 / (rel_s + 1e-9))
        
        current_level = 0.0
        hold_counter = 0
        
        for i in range(len(envelope)):
            target = envelope[i]
            if target > current_level:
                current_level += (target - current_level) * g_att
                hold_counter = hold_s
            elif hold_counter > 0:
                hold_counter -= 1
            else:
                current_level += (target - current_level) * g_rel
            smoothed_env[i] = current_level
            
        envelope = smoothed_env
        if envelope.max() > 0: envelope = envelope / envelope.max()
        duck_curve = 1.0 - (envelope * strength)
        
        # Match length
        target_len = bass_array.shape[0]
        if len(duck_curve) < target_len:
            duck_curve = np.pad(duck_curve, (0, target_len - len(duck_curve)), mode='edge')
        else: duck_curve = duck_curve[:target_len]
            
        if bass_array.ndim > 1:
            duck_curve = np.column_stack((duck_curve, duck_curve))

        if mode.lower() == "multiband":
            sos_lp = butter(4, 200, 'lp', fs=sr, output='sos')
            sos_hp = butter(4, 200, 'hp', fs=sr, output='sos')
            
            lows = np.zeros_like(bass_array); highs = np.zeros_like(bass_array)
            if bass_array.ndim == 1:
                lows = sosfilt(sos_lp, bass_array); highs = sosfilt(sos_hp, bass_array)
            else:
                for ch in range(bass_array.shape[1]):
                    lows[:, ch] = sosfilt(sos_lp, bass_array[:, ch])
                    highs[:, ch] = sosfilt(sos_hp, bass_array[:, ch])
            return (lows * duck_curve) + highs
        else:
            return bass_array * duck_curve