"""Spectral analyzer for AI MusiMuse.

This module defines :class:`SpectralAnalyzer`, which extracts
fundamental spectral features from decoded audio using shared DSP
utilities.

No rhythm, harmony, embeddings, or machine learning is performed.
The analyzer must not perform logging — the pipeline owns all logging.
The analyzer must not access the database — persistence is the
pipeline's responsibility.
"""

from __future__ import annotations

import time

import numpy as np

from analyzer.analyzer import Analyzer
from analyzer.context import AnalysisContext
from analyzer.feature import Feature, FeatureSet
from analyzer.result import AnalysisResult
from core.exceptions import AnalysisError
from dsp.config import FFT_SIZE, HOP_LENGTH
from dsp.spectrum import (
    compute_magnitude_spectrum,
    compute_power_spectrum,
    compute_stft,
    hann_window,
)


def _spectral_centroid(
    magnitude: np.ndarray,
    frequencies: np.ndarray,
) -> float:
    """Compute the spectral centroid (weighted mean frequency).

    Args:
        magnitude: Magnitude spectrum (n_frames, n_bins).
        frequencies: Frequency bin centers (n_bins,).

    Returns:
        Average spectral centroid across all frames in Hz.
    """
    frame_centroids = np.empty(magnitude.shape[0])
    for i in range(magnitude.shape[0]):
        mag = magnitude[i]
        total = np.sum(mag)
        if total <= 0.0:
            frame_centroids[i] = 0.0
        else:
            frame_centroids[i] = np.sum(frequencies * mag) / total
    return float(np.mean(frame_centroids))


def _spectral_bandwidth(
    magnitude: np.ndarray,
    frequencies: np.ndarray,
    centroid: float,
) -> float:
    """Compute the spectral bandwidth (weighted std dev around centroid).

    Args:
        magnitude: Magnitude spectrum (n_frames, n_bins).
        frequencies: Frequency bin centers (n_bins,).
        centroid: Pre-computed spectral centroid.

    Returns:
        Average spectral bandwidth across all frames in Hz.
    """
    frame_bw = np.empty(magnitude.shape[0])
    for i in range(magnitude.shape[0]):
        mag = magnitude[i]
        total = np.sum(mag)
        if total <= 0.0:
            frame_bw[i] = 0.0
        else:
            frame_bw[i] = np.sqrt(np.sum(((frequencies - centroid) ** 2) * mag) / total)
    return float(np.mean(frame_bw))


def _spectral_rolloff(
    magnitude: np.ndarray,
    frequencies: np.ndarray,
    rolloff_percent: float = 0.85,
) -> float:
    """Compute the spectral rolloff frequency.

    The frequency below which ``rolloff_percent`` of cumulative spectral
    energy exists.

    Args:
        magnitude: Magnitude spectrum (n_frames, n_bins).
        frequencies: Frequency bin centers (n_bins,).
        rolloff_percent: Fraction of energy threshold (default 0.85).

    Returns:
        Average rolloff frequency across all frames in Hz.
    """
    frame_rolloff = np.empty(magnitude.shape[0])
    for i in range(magnitude.shape[0]):
        mag = magnitude[i]
        total = np.sum(mag)
        if total <= 0.0:
            frame_rolloff[i] = 0.0
        else:
            cumulative = np.cumsum(mag)
            threshold = total * rolloff_percent
            idx = np.searchsorted(cumulative, threshold)
            idx = min(idx, len(frequencies) - 1)
            frame_rolloff[i] = frequencies[idx]
    return float(np.mean(frame_rolloff))


def _spectral_flatness(power: np.ndarray) -> float:
    """Compute the spectral flatness (Wiener entropy).

    Geometric mean / arithmetic mean of the power spectrum.
    Returns a value in [0, 1] where 1.0 indicates white noise.

    Args:
        power: Power spectrum (n_frames, n_bins).

    Returns:
        Average spectral flatness across all frames.
    """
    frame_flatness = np.empty(power.shape[0])
    for i in range(power.shape[0]):
        p = power[i]
        if np.any(p <= 0.0):
            frame_flatness[i] = 0.0
        else:
            n = len(p)
            log_mean = np.sum(np.log(p)) / n
            geo_mean = np.exp(log_mean)
            arith_mean = np.sum(p) / n
            if arith_mean <= 0.0:
                frame_flatness[i] = 0.0
            else:
                frame_flatness[i] = geo_mean / arith_mean
    return float(np.mean(frame_flatness))


def _zero_crossing_rate(samples: np.ndarray) -> float:
    """Compute the zero crossing rate of a 1-D signal.

    Args:
        samples: 1-D mono audio signal.

    Returns:
        Fraction of consecutive sample pairs that cross zero.
    """
    if len(samples) < 2:
        return 0.0
    signs = np.sign(samples)
    crossings = np.sum(np.abs(np.diff(signs)) > 0)
    return float(crossings) / float(len(samples) - 1)


def _spectral_flux(magnitude: np.ndarray) -> float:
    """Compute the average normalized spectral flux.

    Frame-to-frame difference of magnitude spectra, normalized by
    frame energy.

    Args:
        magnitude: Magnitude spectrum (n_frames, n_bins).

    Returns:
        Average spectral flux across all frames.
    """
    if magnitude.shape[0] < 2:
        return 0.0
    diff = np.diff(magnitude, axis=0)
    flux = np.sqrt(np.sum(diff**2, axis=1))
    return float(np.mean(flux))


class SpectralAnalyzer(Analyzer):
    """Extracts spectral features from decoded audio.

    Produces 6 features: spectral centroid, bandwidth, rolloff,
    flatness, zero crossing rate, and spectral flux.

    All calculations use shared DSP utilities from :mod:`dsp.spectrum`
    and configuration from :mod:`dsp.config`.
    """

    name = "spectral"
    version = "1.0.0"

    def analyze(self, context: AnalysisContext) -> AnalysisResult:
        """Analyze decoded audio and return spectral features.

        Args:
            context: The analysis context containing decoded audio.

        Returns:
            An :class:`AnalysisResult` with spectral feature data.

        Raises:
            AnalysisError: If the audio is empty, contains NaN, or
                contains infinite values.
        """
        start = time.perf_counter()

        audio = context.decoded_audio
        samples = audio.samples

        if samples.size == 0:
            raise AnalysisError("Audio contains no samples")

        if np.any(np.isnan(samples)):
            raise AnalysisError("Audio contains NaN values")

        if np.any(np.isinf(samples)):
            raise AnalysisError("Audio contains infinite values")

        # Convert to mono
        mono = np.mean(samples[:, :2], axis=1) if audio.channels > 1 else samples[:, 0]

        sample_rate = audio.sample_rate
        frequencies = np.fft.rfftfreq(FFT_SIZE, d=1.0 / sample_rate)

        window = hann_window(FFT_SIZE)
        stft = compute_stft(mono, FFT_SIZE, HOP_LENGTH, window)
        magnitude = compute_magnitude_spectrum(stft)
        power = compute_power_spectrum(magnitude)

        centroid = _spectral_centroid(magnitude, frequencies)
        bandwidth = _spectral_bandwidth(magnitude, frequencies, centroid)
        rolloff = _spectral_rolloff(magnitude, frequencies)
        flatness = _spectral_flatness(power)
        zcr = _zero_crossing_rate(mono)
        flux = _spectral_flux(magnitude)

        features: list[Feature] = [
            Feature(
                name="spectral.centroid",
                value=centroid,
                unit="Hz",
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="spectral.bandwidth",
                value=bandwidth,
                unit="Hz",
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="spectral.rolloff",
                value=rolloff,
                unit="Hz",
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="spectral.flatness",
                value=flatness,
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="spectral.zero_crossing_rate",
                value=zcr,
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="spectral.flux",
                value=flux,
                analyzer=self.name,
                version=self.version,
            ),
        ]

        feature_set = FeatureSet(features)
        elapsed_ms = (time.perf_counter() - start) * 1000.0

        return AnalysisResult(
            analyzer_name=self.name,
            analyzer_version=self.version,
            execution_time_ms=elapsed_ms,
            success=True,
            warnings=(),
            feature_set=feature_set,
        )
