"""Reusable rhythm DSP utilities for AI MusiMuse.

This module provides shared helper functions for temporal rhythm
analysis.  All rhythm analyzers must use these utilities instead of
implementing onset detection or tempo estimation independently.

The DSP package is the single source of truth for rhythm processing
throughout the project.
"""

from __future__ import annotations

import numpy as np
from scipy.signal import find_peaks

from dsp.spectrum import compute_magnitude_spectrum, compute_stft


def compute_onset_envelope(
    samples: np.ndarray,
    fft_size: int,
    hop_size: int,
    window: np.ndarray,
) -> np.ndarray:
    """Compute the spectral-flux onset envelope of a 1-D signal.

    The onset envelope is derived from positive frame-to-frame
    magnitude spectrum differences (spectral flux), summed across
    frequency bins.

    Args:
        samples: 1-D mono audio signal.
        fft_size: FFT size in samples.
        hop_size: Hop size between frames in samples.
        window: Window function array (length == fft_size).

    Returns:
        1-D array of onset strength values, one per frame.
    """
    stft = compute_stft(samples, fft_size, hop_size, window)
    magnitude = compute_magnitude_spectrum(stft)

    if magnitude.shape[0] < 2:
        return np.zeros(max(1, magnitude.shape[0]), dtype=np.float64)

    diff = np.diff(magnitude, axis=0)
    flux = np.sum(np.maximum(diff, 0.0), axis=1)
    return flux.astype(np.float64)


def compute_autocorrelation(signal: np.ndarray) -> np.ndarray:
    """Compute the autocorrelation of a 1-D signal.

    Uses FFT-based autocorrelation for efficiency.

    Args:
        signal: 1-D input signal.

    Returns:
        1-D autocorrelation array (same length as input), with zero
        lag at index 0.
    """
    n = len(signal)
    if n == 0:
        return np.array([], dtype=np.float64)

    fft_size = 1
    while fft_size < 2 * n:
        fft_size *= 2

    fft = np.fft.rfft(signal, n=fft_size)
    power = fft * np.conj(fft)
    acf = np.fft.irfft(power, n=fft_size)[:n]
    if acf[0] > 0.0:
        acf = acf / acf[0]
    return acf.astype(np.float64)


def estimate_tempo(
    onset_env: np.ndarray,
    sample_rate: int,
    hop_size: int,
    min_bpm: int,
    max_bpm: int,
) -> tuple[float, float]:
    """Estimate tempo from the onset envelope via autocorrelation.

    Searches for the strongest autocorrelation peak in the lag range
    corresponding to [min_bpm, max_bpm].

    Args:
        onset_env: Onset envelope array.
        sample_rate: Audio sample rate in Hz.
        hop_size: Hop size used to compute the onset envelope.
        min_bpm: Minimum tempo to search (BPM).
        max_bpm: Maximum tempo to search (BPM).

    Returns:
        A tuple of (tempo_bpm, beat_strength) where beat_strength is
        the normalized autocorrelation peak height in [0, 1].
        Returns (0.0, 0.0) for silent or constant signals.
    """
    if len(onset_env) < 2 or np.all(onset_env == 0.0):
        return 0.0, 0.0

    acf = compute_autocorrelation(onset_env)

    frames_per_second = sample_rate / hop_size
    min_lag = int(frames_per_second * 60.0 / max_bpm)
    max_lag = int(frames_per_second * 60.0 / min_bpm)
    max_lag = min(max_lag, len(acf) - 1)

    if min_lag >= max_lag or min_lag >= len(acf):
        return 0.0, 0.0

    search_region = acf[min_lag : max_lag + 1]
    if len(search_region) == 0 or np.all(search_region <= 0.0):
        return 0.0, 0.0

    peak_idx = int(np.argmax(search_region)) + min_lag
    peak_value = float(acf[peak_idx])

    if peak_value <= 0.0:
        return 0.0, 0.0

    tempo_bpm = 60.0 * frames_per_second / peak_idx
    beat_strength = min(max(peak_value, 0.0), 1.0)

    return float(tempo_bpm), beat_strength


def estimate_beat_period(tempo_bpm: float) -> float:
    """Compute the beat period in seconds from tempo.

    Args:
        tempo_bpm: Tempo in beats per minute.

    Returns:
        Beat period in seconds, or 0.0 for zero tempo.
    """
    if tempo_bpm <= 0.0:
        return 0.0
    return 60.0 / tempo_bpm


def compute_rhythm_density(
    onset_env: np.ndarray,
    sample_rate: int,
    hop_size: int,
) -> float:
    """Compute rhythm density (estimated onsets per second).

    Peak-picks the onset envelope and divides the count by the
    total duration.

    Args:
        onset_env: Onset envelope array.
        sample_rate: Audio sample rate in Hz.
        hop_size: Hop size used to compute the onset envelope.

    Returns:
        Estimated onsets per second.
    """
    if len(onset_env) == 0 or np.all(onset_env == 0.0):
        return 0.0

    threshold = 0.1 * float(np.max(onset_env))
    peaks, _ = find_peaks(onset_env, height=threshold, distance=1)

    duration_seconds = len(onset_env) * hop_size / sample_rate
    if duration_seconds <= 0.0:
        return 0.0

    return float(len(peaks)) / duration_seconds


def detect_onsets(
    onset_env: np.ndarray,
    sample_rate: int,
    hop_size: int,
) -> tuple[int, float, float]:
    """Detect onsets by peak-picking the onset envelope.

    Args:
        onset_env: Onset envelope array.
        sample_rate: Audio sample rate in Hz.
        hop_size: Hop size used to compute the onset envelope.

    Returns:
        A tuple of (onset_count, onset_rate, rhythm_regularity)
        where rhythm_regularity is in [0, 1].
    """
    if len(onset_env) == 0 or np.all(onset_env == 0.0):
        return 0, 0.0, 0.0

    threshold = 0.1 * float(np.max(onset_env))
    peaks, _ = find_peaks(onset_env, height=threshold, distance=1)

    onset_count = len(peaks)
    duration_seconds = len(onset_env) * hop_size / sample_rate
    onset_rate = (
        float(onset_count) / duration_seconds if duration_seconds > 0.0 else 0.0
    )

    if onset_count < 2:
        regularity = 0.0
    else:
        intervals = np.diff(peaks)
        if len(intervals) == 0:
            regularity = 0.0
        else:
            mean_interval = float(np.mean(intervals))
            if mean_interval <= 0.0:
                regularity = 0.0
            else:
                std_interval = float(np.std(intervals))
                cv = std_interval / mean_interval
                regularity = max(0.0, min(1.0, 1.0 / (1.0 + cv)))

    return onset_count, onset_rate, regularity
