"""Reusable harmonic DSP utilities for AI MusiMuse.

This module provides shared helper functions for harmonic analysis.
All harmony analyzers must use these utilities instead of implementing
chroma or key estimation independently.

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

from __future__ import annotations

import numpy as np

from dsp.spectrum import (
    compute_magnitude_spectrum,
    compute_power_spectrum,
    compute_stft,
)

# Reference frequency for C0 (MIDI note 0 is C-1, but we use C0 = 16.035 Hz)
# Actually, standard: A4 = 440 Hz, C0 = 440 * 2^(-4.75) ≈ 16.35 Hz
F_REF = 440.0 * 2.0 ** (-4.75)

# Key names matching the task spec (mixed sharps/flats)
KEY_NAMES = [
    "C",
    "C#",
    "D",
    "Eb",
    "E",
    "F",
    "F#",
    "G",
    "Ab",
    "A",
    "Bb",
    "B",
]

# Krumhansl-Schmuckler key profiles
MAJOR_PROFILE = np.array(
    [6.35, 2.23, 3.48, 2.33, 4.38, 4.09, 2.52, 5.19, 2.39, 3.66, 2.29, 2.88],
    dtype=np.float64,
)
MINOR_PROFILE = np.array(
    [6.33, 2.68, 3.52, 5.38, 2.60, 3.53, 2.54, 4.75, 3.98, 2.69, 3.34, 3.17],
    dtype=np.float64,
)

# Tonnetz transformation matrix (6 x 12)
# Maps 12 pitch classes to 6-dimensional tonal space:
# (fifths x/y, minor thirds x/y, major thirds x/y)
_FIFTHS = np.arange(12, dtype=np.float64) * 7.0  # circle of fifths
_MINOR_THIRDS = np.arange(12, dtype=np.float64) * 3.0  # circle of minor thirds
_MAJOR_THIRDS = np.arange(12, dtype=np.float64) * 4.0  # circle of major thirds

_FIFTHS_ANGLE = _FIFTHS * (2.0 * np.pi / 12.0)
_MINOR_ANGLE = _MINOR_THIRDS * (2.0 * np.pi / 12.0)
_MAJOR_ANGLE = _MAJOR_THIRDS * (2.0 * np.pi / 12.0)

TONNETZ_MATRIX = np.array(
    [
        np.cos(_FIFTHS_ANGLE),
        np.sin(_FIFTHS_ANGLE),
        np.cos(_MINOR_ANGLE),
        np.sin(_MINOR_ANGLE),
        np.cos(_MAJOR_ANGLE),
        np.sin(_MAJOR_ANGLE),
    ],
    dtype=np.float64,
)


def compute_chroma(
    samples: np.ndarray,
    fft_size: int,
    hop_size: int,
    window: np.ndarray,
    sample_rate: int,
    n_bins: int = 12,
) -> np.ndarray:
    """Compute a normalized 12-bin chroma vector from audio.

    Maps spectral frequency bins to pitch classes using
    ``12 * log2(f / f_ref) mod 12``, sums power per class, and
    normalizes to [0, 1].

    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).
        sample_rate: Audio sample rate in Hz.
        n_bins: Number of chroma bins (default 12).

    Returns:
        Normalized chroma vector of length ``n_bins``.
    """
    stft = compute_stft(samples, fft_size, hop_size, window)
    magnitude = compute_magnitude_spectrum(stft)
    power = compute_power_spectrum(magnitude)

    n_frames, n_freq_bins = power.shape
    freqs = np.fft.rfftfreq(fft_size, d=1.0 / sample_rate)

    # Skip DC bin (index 0)
    valid = freqs[1:] > 0.0
    freqs_valid = freqs[1:][valid]
    power_valid = power[:, 1:][:, valid]

    if len(freqs_valid) == 0:
        return np.zeros(n_bins, dtype=np.float64)

    # Map frequencies to pitch classes
    pitch_classes = (np.round(12.0 * np.log2(freqs_valid / F_REF)).astype(int)) % n_bins

    # Aggregate power per pitch class across all frames
    chroma = np.zeros(n_bins, dtype=np.float64)
    for pc in range(n_bins):
        mask = pitch_classes == pc
        if np.any(mask):
            chroma[pc] = float(np.sum(power_valid[:, mask]))

    # Normalize
    total = float(np.sum(chroma))
    if total > 0.0:
        chroma = chroma / total

    return chroma


def compute_pitch_class_histogram(chroma: np.ndarray) -> np.ndarray:
    """Return the raw pitch class histogram from a chroma vector.

    This is identical to the chroma vector but provided as a separate
    function for semantic clarity in future use cases.

    Args:
        chroma: Chroma vector (normalized or unnormalized).

    Returns:
        The same vector as a histogram.
    """
    return chroma.copy()


def estimate_key(chroma: np.ndarray) -> str:
    """Estimate the musical key using the Krumhansl-Schmuckler algorithm.

    Correlates the chroma vector with all 12 rotations of the major
    and minor key profiles.  Returns the key name corresponding to
    the best correlation.

    Args:
        chroma: 12-bin chroma vector.

    Returns:
        Key name string (e.g. "C", "C#", "D", "Eb", ...).
    """
    if np.all(chroma == 0.0):
        return "C"

    best_score = -np.inf
    best_key_idx = 0

    for i in range(12):
        rotated_major = np.roll(MAJOR_PROFILE, i)
        rotated_minor = np.roll(MINOR_PROFILE, i)

        score_major = float(np.corrcoef(chroma, rotated_major)[0, 1])
        score_minor = float(np.corrcoef(chroma, rotated_minor)[0, 1])

        score = max(score_major, score_minor)
        if score > best_score:
            best_score = score
            best_key_idx = i

    return KEY_NAMES[best_key_idx]


def estimate_mode(chroma: np.ndarray) -> str:
    """Estimate the musical mode (major or minor).

    Compares the best major profile correlation against the best
    minor profile correlation.

    Args:
        chroma: 12-bin chroma vector.

    Returns:
        "major", "minor", or "unknown" for silence.
    """
    if np.all(chroma == 0.0):
        return "unknown"

    best_major = -np.inf
    best_minor = -np.inf

    for i in range(12):
        rotated_major = np.roll(MAJOR_PROFILE, i)
        rotated_minor = np.roll(MINOR_PROFILE, i)

        score_major = float(np.corrcoef(chroma, rotated_major)[0, 1])
        score_minor = float(np.corrcoef(chroma, rotated_minor)[0, 1])

        best_major = max(best_major, score_major)
        best_minor = max(best_minor, score_minor)

    return "major" if best_major >= best_minor else "minor"


def compute_tonal_centroid(chroma: np.ndarray) -> np.ndarray:
    """Compute the 6-dimensional tonal centroid (tonnetz).

    Transforms the chroma vector into the tonal space using the
    tonnetz transformation matrix.

    Args:
        chroma: 12-bin chroma vector.

    Returns:
        6-element array [x, y, z, u, v, w].
    """
    if np.all(chroma == 0.0):
        return np.zeros(6, dtype=np.float64)

    total = float(np.sum(chroma))
    if total <= 0.0:
        return np.zeros(6, dtype=np.float64)

    centroid = TONNETZ_MATRIX @ chroma
    return centroid.astype(np.float64)


def compute_tonal_stability(chroma: np.ndarray) -> float:
    """Compute tonal stability from the chroma vector.

    Defined as ``1.0 - normalized_entropy`` where normalized entropy
    is the Shannon entropy of the chroma distribution divided by the
    maximum possible entropy (log2(12)).

    Args:
        chroma: 12-bin chroma vector.

    Returns:
        Tonal stability in [0, 1].  Higher values indicate more
        stable harmonic content.
    """
    if np.all(chroma == 0.0):
        return 0.0

    total = float(np.sum(chroma))
    if total <= 0.0:
        return 0.0

    probs = chroma / total
    probs = probs[probs > 0.0]

    entropy = float(-np.sum(probs * np.log2(probs)))
    max_entropy = float(np.log2(len(chroma)))

    if max_entropy <= 0.0:
        return 0.0

    return float(1.0 - entropy / max_entropy)
