import numpy as np
import os
import soundfile as sf
from scipy.signal import butter, sosfilt

try:
    from .modules.sidechain import SidechainModule
    from .modules.exciter import ExciterModule
    from .modules.transient import TransientModule
    from .modules.dynamics import DynamicsModule
    from .modules.imaging import ImagingModule
    from .modules.vintage import VintageModule
    from .modules.harmonics import HarmonicsModule
    from .modules.eq import EqModule
    from .modules.ms_limiter import MsLimiterModule
    from .modules.spectral_smoother import SpectralSmootherModule
    from .modules.sub_synth import SubSynthModule
    from .pipeline.tp_limiter import limit
except (ImportError, ValueError):
    from modules.sidechain import SidechainModule
    from modules.exciter import ExciterModule
    from modules.transient import TransientModule
    from modules.dynamics import DynamicsModule
    from modules.imaging import ImagingModule
    from modules.vintage import VintageModule
    from modules.harmonics import HarmonicsModule
    from modules.eq import EqModule
    from modules.ms_limiter import MsLimiterModule
    from modules.spectral_smoother import SpectralSmootherModule
    from modules.sub_synth import SubSynthModule
    from automaster_app.pipeline.tp_limiter import limit


def _legacy_limiter_enabled(use_legacy_limiter=None):
    if use_legacy_limiter is not None:
        return use_legacy_limiter
    return os.environ.get("AUTOREMASTER_LEGACY_CHAIN", "").lower() in ("1", "true", "yes")


class AudioProcessorPro:
    """Pro Audio Processor for AutoRemaster CLI."""
    
    def __init__(self, use_legacy_limiter=None):
        self.sample_rate = 44100
        self.use_legacy_limiter = _legacy_limiter_enabled(use_legacy_limiter)
        self.sc_module = None
        self.exciter_module = None
        self.transient_module = None
        self.dyn_module = None
        self.img_module = None
        self.vintage_module = None
        self.harm_module = None
        self.eq_module = None
        self.ms_limiter = None
        self.smoother = None
        self.sub_synth = None

    def load_file(self, path):
        data, sr = sf.read(path)
        self.sample_rate = sr
        if data.ndim == 1: data = np.column_stack((data, data))
        
        # Init modules
        self.sc_module = SidechainModule(sr)
        self.exciter_module = ExciterModule(sr)
        self.transient_module = TransientModule(sr)
        self.dyn_module = DynamicsModule(sr)
        self.img_module = ImagingModule(sr)
        self.vintage_module = VintageModule(sr)
        self.harm_module = HarmonicsModule(sr)
        self.eq_module = EqModule(sr)
        self.ms_limiter = MsLimiterModule(sr)
        self.smoother = SpectralSmootherModule(sr)
        self.sub_synth = SubSynthModule(sr)
        
        return data, sr

    def process_chain(self, audio, preset, log_callback=None, progress_callback=None):
        """Execute the mastering chain."""
        order = preset.get("module_order", [])
        active = preset.get("active_modules", {})
        
        # Calculate progress increments
        active_keys = [k for k in order if active.get(k)]
        num_modules = len(active_keys)
        
        # We start at 10% (after loading) and end at 95% (before saving)
        start_pct = 10
        total_range = 85
        increment = total_range / num_modules if num_modules > 0 else 0
        current_progress = start_pct

        for idx, key in enumerate(active_keys):
            if log_callback: log_callback(f"Applying {key.replace('_', ' ').title()}...")
            
            if key == "ms_limiter":
                audio = self.ms_limiter.process(
                    audio, 
                    side_limit_db=preset.get("ms_side_limit", -3.0),
                    mid_gain_db=preset.get("ms_mid_gain", 0.0)
                )

            elif key == "sub_synth":
                audio = self.sub_synth.process(
                    audio,
                    amount=preset.get("sub_synth_amount", 0.0),
                    freq_hz=preset.get("sub_synth_freq", 60.0)
                )

            elif key == "mono_bass":
                audio = self.img_module.make_mono_bass(audio, preset.get("mono_cutoff_hz", 120))
                
            elif key == "transient_shaper":
                audio = self.transient_module.process(
                    audio, 
                    preset.get("transient_boost_db", 0), 
                    -2.0
                )
                
            elif key == "auto_sidechain":
                audio = self.sc_module.process(
                    audio, audio, 
                    strength=preset.get("sidechain_strength", 0),
                    mode=preset.get("sidechain_mode", "broadband"),
                    hold_ms=preset.get("sidechain_hold_ms", 10.0),
                    release_ms=preset.get("sidechain_release_ms", 100.0)
                )
                
            elif key == "dynamic_eq":
                audio = self.eq_module.apply_dynamic_eq(audio, preset.get("dynamic_eq_settings", []))
                
            elif key == "multiband":
                audio = self.dyn_module.apply_multiband(audio, preset.get("multiband_settings", []))
                
            elif key == "vintage":
                audio = self.vintage_module.process(
                    audio,
                    preset.get("vintage_saturation", 0),
                    preset.get("harmonic_character", "Tube"),
                    blend=preset.get("vintage_blend", 1.0)
                )
                
            elif key == "harmonic_enhancement":
                audio = self.harm_module.process(
                    audio,
                    preset.get("warmth_amount", 0),
                    preset.get("edge_amount", 0),
                    blend=preset.get("harmonics_blend", 1.0)
                )
                
            elif key == "ai_dehaze": # Exciter
                amt = preset.get("exciter_amount", "Med")
                if amt != "None":
                    audio = self.exciter_module.process(audio, mode="spectral", blend_db=-12)
                    
            elif key == "hf_restoration":
                if preset.get("hf_restoration_mode") == "dsp_fast":
                    audio = self.exciter_module.process(audio, mode="dsp_fast")
            
            elif key == "spectral_smoother":
                # Find the base frequency of the key for anti-aliasing
                key_name = preset.get("key", "C Major")
                # Approximate base freq (C1 range)
                base_freqs = {'C': 32.7, 'C#': 34.6, 'D': 36.7, 'D#': 38.9, 'E': 41.2, 'F': 43.7, 'F#': 46.2, 'G': 49.0, 'G#': 51.9, 'A': 55.0, 'A#': 58.3, 'B': 61.7}
                k_note = str(key_name).split(' ')[0]
                k_freq = base_freqs.get(k_note, 44.0)

                audio = self.smoother.process(
                    audio,
                    resonances=preset.get("resonances", []),
                    threshold_db=preset.get("smoother_threshold", -20.0),
                    sensitivity=preset.get("smoother_sensitivity", 0.5),
                    key_freq=k_freq
                )
                    
            elif key == "stereo_imaging":
                # 1. Center the image
                audio = self.img_module.apply_balancing(
                    audio, 
                    balance_db=preset.get("image_balance_db", 0.0)
                )
                
                # 2. Synthesize Width (for mono sources)
                audio = self.img_module.apply_haas_widening(
                    audio,
                    amount=preset.get("haas_amount", 0.0)
                )
                
                # 3. Apply standard Width
                audio = self.img_module.apply_width(
                    audio,
                    mid_w=preset.get("stereo_mid_width", 1.0),
                    high_w=preset.get("stereo_high_width", 1.0)
                )
                
                # --- SPECTRAL SILK (Final Anti-Aliasing) ---
                # Apply ultra-gentle low pass at 19kHz to remove digital 'shards'
                sos_silk = butter(4, 19000, 'lp', fs=self.sample_rate, output='sos')
                if audio.ndim == 1: audio = sosfilt(sos_silk, audio)
                else:
                    audio[:, 0] = sosfilt(sos_silk, audio[:, 0])
                    audio[:, 1] = sosfilt(sos_silk, audio[:, 1])
                
            elif key == "final_limiting":
                target_lufs = preset.get("target_lufs", -8.0)
                if self.use_legacy_limiter:
                    audio = self.dyn_module.apply_limiter(
                        audio,
                        target_lufs,
                        release_ms=preset.get("limiter_release_ms", 50.0),
                    )
                else:
                    audio, _ = limit(
                        audio,
                        self.sample_rate,
                        target_lufs=target_lufs,
                        ceiling_dbtp=-1.0,
                    )
            
            # Update incremental progress
            current_progress += increment
            if progress_callback:
                progress_callback(int(current_progress))
                
        return audio

    def save(self, audio, path):
        sf.write(path, audio, self.sample_rate, subtype='PCM_24')
