"""Rhythm analyzer for AI MusiMuse.

This module defines :class:`RhythmAnalyzer`, which extracts temporal
rhythmic characteristics from decoded audio using shared DSP utilities.

No beat tracking, beat alignment, downbeat detection, swing estimation,
or groove analysis 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, ONSET_HOP_SIZE, TEMPO_MAX_BPM, TEMPO_MIN_BPM
from dsp.rhythm import (
    compute_onset_envelope,
    compute_rhythm_density,
    estimate_beat_period,
    estimate_tempo,
)
from dsp.spectrum import hann_window


class RhythmAnalyzer(Analyzer):
    """Extracts rhythm features from decoded audio.

    Produces 7 features: tempo, beat period, rhythm density, beat
    strength, onset count, onset rate, and rhythm regularity.

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

    name = "rhythm"
    version = "1.0.0"

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

        Args:
            context: The analysis context containing decoded audio.

        Returns:
            An :class:`AnalysisResult` with rhythm 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

        # Onset envelope via spectral flux
        window = hann_window(FFT_SIZE)
        onset_env = compute_onset_envelope(mono, FFT_SIZE, ONSET_HOP_SIZE, window)

        # Tempo and beat strength from autocorrelation
        tempo_bpm, beat_strength = estimate_tempo(
            onset_env,
            sample_rate,
            ONSET_HOP_SIZE,
            TEMPO_MIN_BPM,
            TEMPO_MAX_BPM,
        )

        # Beat period
        beat_period = estimate_beat_period(tempo_bpm)

        # Rhythm density
        rhythm_density = compute_rhythm_density(onset_env, sample_rate, ONSET_HOP_SIZE)

        # Onset detection
        from dsp.rhythm import detect_onsets

        onset_count, onset_rate, rhythm_regularity = detect_onsets(
            onset_env, sample_rate, ONSET_HOP_SIZE
        )

        features: list[Feature] = [
            Feature(
                name="rhythm.tempo",
                value=tempo_bpm,
                unit="BPM",
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="rhythm.beat_period",
                value=beat_period,
                unit="seconds",
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="rhythm.density",
                value=rhythm_density,
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="rhythm.beat_strength",
                value=beat_strength,
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="rhythm.onset_count",
                value=float(onset_count),
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="rhythm.onset_rate",
                value=onset_rate,
                analyzer=self.name,
                version=self.version,
            ),
            Feature(
                name="rhythm.regularity",
                value=rhythm_regularity,
                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,
        )
