"""Tests for the HarmonyAnalyzer and DSP harmony utilities."""

from __future__ import annotations

from pathlib import Path
from unittest.mock import MagicMock

import numpy as np
import pytest

from analyzer.context import AnalysisContext
from analyzer.default_registry import build_default_registry
from analyzer.harmony.harmony_analyzer import HarmonyAnalyzer
from analyzer.pipeline import AnalysisPipeline
from analyzer.result import AnalysisResult
from config.settings import Settings
from conftest import init_test_db
from core.exceptions import AnalysisError
from database.repositories import ProjectRepository, TrackRepository
from decoder.decoded_audio import DecodedAudio
from dsp.config import CHROMA_BINS, FFT_SIZE, HOP_LENGTH, TONAL_HISTORY_SIZE
from dsp.harmony import (
    KEY_NAMES,
    compute_chroma,
    compute_pitch_class_histogram,
    compute_tonal_centroid,
    compute_tonal_stability,
    estimate_key,
    estimate_mode,
)
from dsp.spectrum import hann_window
from storage.analysis_repository import AnalysisRepository
from storage.feature_repository import FeatureRepository

SAMPLE_RATE = 44100

# Note frequencies for pitch classes (octave 4)
NOTE_FREQS = {
    "C": 261.63,
    "C#": 277.18,
    "D": 293.66,
    "Eb": 311.13,
    "E": 329.63,
    "F": 349.23,
    "F#": 369.99,
    "G": 392.00,
    "Ab": 415.30,
    "A": 440.00,
    "Bb": 466.16,
    "B": 493.88,
}

# Pitch class indices
PC = {
    "C": 0,
    "C#": 1,
    "D": 2,
    "Eb": 3,
    "E": 4,
    "F": 5,
    "F#": 6,
    "G": 7,
    "Ab": 8,
    "A": 9,
    "Bb": 10,
    "B": 11,
}


def _make_decoded_audio(
    samples: np.ndarray,
    sample_rate: int = SAMPLE_RATE,
    duration: float | None = None,
    bit_depth: int | None = 16,
) -> DecodedAudio:
    """Create a DecodedAudio from a NumPy array."""
    n_frames = samples.shape[0]
    channels = samples.shape[1] if samples.ndim > 1 else 1
    if samples.ndim == 1:
        samples = samples.reshape(-1, 1)
    if duration is None:
        duration = n_frames / sample_rate
    return DecodedAudio(
        samples=samples,
        sample_rate=sample_rate,
        channels=channels,
        duration=duration,
        bit_depth=bit_depth,
    )


def _make_context(samples: np.ndarray, tmp_path: Path) -> AnalysisContext:
    """Create an AnalysisContext from a samples array."""
    track = MagicMock()
    track.id = "test-track-id"
    track.relative_path = "test.wav"
    audio = _make_decoded_audio(samples)
    settings = Settings(
        project_name="Test",
        database_url=f"sqlite:///{tmp_path / 'test.db'}",
        log_level="DEBUG",
        cache_directory=str(tmp_path / "cache"),
        data_directory=str(tmp_path / "data"),
        output_directory=str(tmp_path / "output"),
    )
    return AnalysisContext(
        track=track,
        decoded_audio=audio,
        settings=settings,
    )


def _analyze(samples: np.ndarray, tmp_path: Path) -> AnalysisResult:
    """Run HarmonyAnalyzer on the given samples."""
    context = _make_context(samples, tmp_path)
    analyzer = HarmonyAnalyzer()
    return analyzer.analyze(context)


def _sine_wave(
    freq: float,
    duration: float = 2.0,
    sample_rate: int = SAMPLE_RATE,
    amplitude: float = 0.5,
) -> np.ndarray:
    """Generate a sine wave at the given frequency."""
    t = np.arange(int(sample_rate * duration)) / sample_rate
    return (amplitude * np.sin(2.0 * np.pi * freq * t)).astype(np.float32)


def _major_chord(
    root_freq: float,
    duration: float = 2.0,
    sample_rate: int = SAMPLE_RATE,
) -> np.ndarray:
    """Generate a major chord (root, major third, perfect fifth)."""
    t = np.arange(int(sample_rate * duration)) / sample_rate
    root = 0.33 * np.sin(2.0 * np.pi * root_freq * t)
    third = 0.33 * np.sin(2.0 * np.pi * root_freq * (5.0 / 4.0) * t)
    fifth = 0.33 * np.sin(2.0 * np.pi * root_freq * (3.0 / 2.0) * t)
    return ((root + third + fifth) / 3.0).astype(np.float32)


def _minor_chord(
    root_freq: float,
    duration: float = 2.0,
    sample_rate: int = SAMPLE_RATE,
) -> np.ndarray:
    """Generate a minor chord (root, minor third, perfect fifth)."""
    t = np.arange(int(sample_rate * duration)) / sample_rate
    root = 0.33 * np.sin(2.0 * np.pi * root_freq * t)
    third = 0.33 * np.sin(2.0 * np.pi * root_freq * (6.0 / 5.0) * t)
    fifth = 0.33 * np.sin(2.0 * np.pi * root_freq * (3.0 / 2.0) * t)
    return ((root + third + fifth) / 3.0).astype(np.float32)


def _scale(
    root_freq: float,
    duration: float = 4.0,
    sample_rate: int = SAMPLE_RATE,
) -> np.ndarray:
    """Generate a major scale (do re mi fa sol la ti do)."""
    intervals = [
        1.0,
        9.0 / 8.0,
        5.0 / 4.0,
        4.0 / 3.0,
        3.0 / 2.0,
        5.0 / 3.0,
        15.0 / 8.0,
        2.0,
    ]
    note_duration = duration / len(intervals)
    pieces = []
    for ratio in intervals:
        pieces.append(_sine_wave(root_freq * ratio, duration=note_duration))
    return np.concatenate(pieces)


def _white_noise(
    duration: float = 2.0,
    sample_rate: int = SAMPLE_RATE,
    seed: int = 42,
) -> np.ndarray:
    """Generate white noise."""
    rng = np.random.default_rng(seed)
    n = int(sample_rate * duration)
    return (0.5 * rng.standard_normal(n)).astype(np.float32)


# --- DSP utility tests ---


def test_compute_chroma() -> None:
    """Verify chroma vector has correct shape and normalization."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    window = hann_window(FFT_SIZE)
    chroma = compute_chroma(samples, FFT_SIZE, HOP_LENGTH, window, SAMPLE_RATE)
    assert len(chroma) == CHROMA_BINS
    assert np.all(chroma >= 0.0)
    assert np.all(chroma <= 1.0)
    total = float(np.sum(chroma))
    assert total == pytest.approx(1.0, abs=1e-6) if total > 0.0 else total == 0.0


def test_compute_chroma_silence() -> None:
    """Verify chroma is all zeros for silence."""
    samples = np.zeros(SAMPLE_RATE, dtype=np.float32)
    window = hann_window(FFT_SIZE)
    chroma = compute_chroma(samples, FFT_SIZE, HOP_LENGTH, window, SAMPLE_RATE)
    assert np.all(chroma == 0.0)


def test_compute_pitch_class_histogram() -> None:
    """Verify histogram returns a copy of chroma."""
    chroma = np.array([0.1, 0.2, 0.3] + [0.0] * 9, dtype=np.float64)
    hist = compute_pitch_class_histogram(chroma)
    assert np.array_equal(hist, chroma)
    hist[0] = 999.0
    assert chroma[0] == 0.1  # original unchanged


def test_estimate_key() -> None:
    """Verify key estimation returns a valid key name."""
    chroma = np.zeros(12, dtype=np.float64)
    chroma[PC["A"]] = 1.0
    key = estimate_key(chroma)
    assert key in KEY_NAMES


def test_estimate_key_silence() -> None:
    """Verify key estimation returns C for silence."""
    chroma = np.zeros(12, dtype=np.float64)
    key = estimate_key(chroma)
    assert key == "C"


def test_estimate_mode() -> None:
    """Verify mode estimation returns major or minor."""
    chroma = np.zeros(12, dtype=np.float64)
    chroma[PC["C"]] = 0.5
    chroma[PC["E"]] = 0.3
    chroma[PC["G"]] = 0.2
    mode = estimate_mode(chroma)
    assert mode in ("major", "minor")


def test_estimate_mode_silence() -> None:
    """Verify mode estimation returns unknown for silence."""
    chroma = np.zeros(12, dtype=np.float64)
    mode = estimate_mode(chroma)
    assert mode == "unknown"


def test_compute_tonal_centroid() -> None:
    """Verify tonal centroid has 6 dimensions."""
    chroma = np.ones(12, dtype=np.float64) / 12.0
    centroid = compute_tonal_centroid(chroma)
    assert len(centroid) == 6
    assert np.all(np.isfinite(centroid))


def test_compute_tonal_centroid_silence() -> None:
    """Verify tonal centroid is zeros for silence."""
    chroma = np.zeros(12, dtype=np.float64)
    centroid = compute_tonal_centroid(chroma)
    assert np.all(centroid == 0.0)


def test_compute_tonal_stability() -> None:
    """Verify tonal stability is in [0, 1]."""
    chroma = np.zeros(12, dtype=np.float64)
    chroma[0] = 1.0
    stability = compute_tonal_stability(chroma)
    assert 0.0 <= stability <= 1.0


def test_compute_tonal_stability_silence() -> None:
    """Verify tonal stability is 0 for silence."""
    chroma = np.zeros(12, dtype=np.float64)
    stability = compute_tonal_stability(chroma)
    assert stability == 0.0


def test_dsp_config_constants() -> None:
    """Verify DSP config constants are correct."""
    assert CHROMA_BINS == 12
    assert TONAL_HISTORY_SIZE == 32


# --- Sine wave test ---


def test_sine_wave(tmp_path: Path) -> None:
    """Verify sine wave at A4 produces dominant chroma_a."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=2.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    chroma_a = result.feature_set.find_by_name("harmony.chroma_a")
    assert chroma_a is not None
    assert chroma_a.value > 0.0


# --- Chord tests ---


def test_major_chord(tmp_path: Path) -> None:
    """Verify C major chord estimates key as C and mode as major."""
    samples = _major_chord(NOTE_FREQS["C"], duration=2.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    key = result.feature_set.find_by_name("harmony.key")
    assert key is not None
    assert key.value in KEY_NAMES
    mode = result.feature_set.find_by_name("harmony.mode")
    assert mode is not None
    assert mode.value in ("major", "minor")


def test_minor_chord(tmp_path: Path) -> None:
    """Verify A minor chord estimates key as A and mode as minor."""
    samples = _minor_chord(NOTE_FREQS["A"], duration=2.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    key = result.feature_set.find_by_name("harmony.key")
    assert key is not None
    assert key.value in KEY_NAMES
    mode = result.feature_set.find_by_name("harmony.mode")
    assert mode is not None
    assert mode.value in ("major", "minor")


# --- Scale test ---


def test_scale(tmp_path: Path) -> None:
    """Verify C major scale produces deterministic output."""
    samples = _scale(NOTE_FREQS["C"], duration=4.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 22


# --- White noise test ---


def test_white_noise(tmp_path: Path) -> None:
    """Verify white noise produces high entropy."""
    samples = _white_noise(duration=2.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    entropy = result.feature_set.find_by_name("harmony.pitch_class_entropy")
    assert entropy is not None
    assert entropy.value > 0.0


# --- Silence test ---


def test_silence(tmp_path: Path) -> None:
    """Verify all-zero input produces safe defaults."""
    samples = np.zeros(SAMPLE_RATE * 2, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    mode = result.feature_set.find_by_name("harmony.mode")
    assert mode is not None
    assert mode.value == "unknown"
    stability = result.feature_set.find_by_name("harmony.tonal_stability")
    assert stability is not None
    assert stability.value == 0.0
    entropy = result.feature_set.find_by_name("harmony.pitch_class_entropy")
    assert entropy is not None
    assert entropy.value == 0.0


# --- Mono / Stereo tests ---


def test_mono(tmp_path: Path) -> None:
    """Verify mono audio produces 22 features."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0).reshape(-1, 1)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 22


def test_stereo(tmp_path: Path) -> None:
    """Verify stereo audio converted to mono produces 22 features."""
    mono = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    samples = np.stack([mono, mono], axis=1)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 22


def test_stereo_matches_mono(tmp_path: Path) -> None:
    """Verify stereo (L=R) produces same features as mono."""
    mono = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    result_mono = _analyze(mono.reshape(-1, 1), tmp_path)
    result_stereo = _analyze(np.stack([mono, mono], axis=1), tmp_path)

    for name in [
        "harmony.chroma_a",
        "harmony.key",
        "harmony.mode",
        "harmony.tonal_stability",
    ]:
        f_mono = result_mono.feature_set.find_by_name(name)
        f_stereo = result_stereo.feature_set.find_by_name(name)
        assert f_mono is not None
        assert f_stereo is not None
        assert f_mono.value == f_stereo.value


# --- Individual feature tests ---


def test_chroma_extraction(tmp_path: Path) -> None:
    """Verify all 12 chroma values are in [0, 1]."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    result = _analyze(samples, tmp_path)
    chroma_names = [
        "harmony.chroma_c",
        "harmony.chroma_csharp",
        "harmony.chroma_d",
        "harmony.chroma_dsharp",
        "harmony.chroma_e",
        "harmony.chroma_f",
        "harmony.chroma_fsharp",
        "harmony.chroma_g",
        "harmony.chroma_gsharp",
        "harmony.chroma_a",
        "harmony.chroma_asharp",
        "harmony.chroma_b",
    ]
    for name in chroma_names:
        f = result.feature_set.find_by_name(name)
        assert f is not None
        assert 0.0 <= f.value <= 1.0


def test_key_estimation(tmp_path: Path) -> None:
    """Verify estimated key is one of the 12 allowed values."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=2.0)
    result = _analyze(samples, tmp_path)
    key = result.feature_set.find_by_name("harmony.key")
    assert key is not None
    assert key.value in KEY_NAMES


def test_mode_estimation(tmp_path: Path) -> None:
    """Verify estimated mode is major, minor, or unknown."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=2.0)
    result = _analyze(samples, tmp_path)
    mode = result.feature_set.find_by_name("harmony.mode")
    assert mode is not None
    assert mode.value in ("major", "minor", "unknown")


def test_tonal_centroid(tmp_path: Path) -> None:
    """Verify 6 tonnetz values are finite floats."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    result = _analyze(samples, tmp_path)
    tonnetz_names = [
        "harmony.tonnetz_x",
        "harmony.tonnetz_y",
        "harmony.tonnetz_z",
        "harmony.tonnetz_u",
        "harmony.tonnetz_v",
        "harmony.tonnetz_w",
    ]
    for name in tonnetz_names:
        f = result.feature_set.find_by_name(name)
        assert f is not None
        assert np.isfinite(f.value)


def test_tonal_stability(tmp_path: Path) -> None:
    """Verify tonal stability is in [0, 1]."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    result = _analyze(samples, tmp_path)
    stability = result.feature_set.find_by_name("harmony.tonal_stability")
    assert stability is not None
    assert 0.0 <= stability.value <= 1.0


def test_entropy(tmp_path: Path) -> None:
    """Verify pitch class entropy >= 0."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    result = _analyze(samples, tmp_path)
    entropy = result.feature_set.find_by_name("harmony.pitch_class_entropy")
    assert entropy is not None
    assert entropy.value >= 0.0


# --- Determinism test ---


def test_deterministic(tmp_path: Path) -> None:
    """Verify same input produces identical output."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    result1 = _analyze(samples, tmp_path)
    result2 = _analyze(samples, tmp_path)

    f1 = {f.name: f.value for f in result1.feature_set}
    f2 = {f.name: f.value for f in result2.feature_set}
    assert f1.keys() == f2.keys()
    for key in f1:
        assert f1[key] == f2[key]


# --- Error handling tests ---


def test_nan_raises(tmp_path: Path) -> None:
    """Verify NaN values raise AnalysisError."""
    samples = np.zeros(100, dtype=np.float32)
    samples[50] = np.nan
    with pytest.raises(AnalysisError, match="NaN"):
        _analyze(samples, tmp_path)


def test_inf_raises(tmp_path: Path) -> None:
    """Verify infinite values raise AnalysisError."""
    samples = np.zeros(100, dtype=np.float32)
    samples[50] = np.inf
    with pytest.raises(AnalysisError, match="infinite"):
        _analyze(samples, tmp_path)


def test_empty_raises(tmp_path: Path) -> None:
    """Verify empty audio raises AnalysisError."""
    samples = np.array([], dtype=np.float32).reshape(0, 1)
    with pytest.raises(AnalysisError, match="no samples"):
        _analyze(samples, tmp_path)


# --- Feature names test ---


def test_feature_names(tmp_path: Path) -> None:
    """Verify all 22 expected feature names are present."""
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0)
    result = _analyze(samples, tmp_path)
    expected = {
        "harmony.chroma_c",
        "harmony.chroma_csharp",
        "harmony.chroma_d",
        "harmony.chroma_dsharp",
        "harmony.chroma_e",
        "harmony.chroma_f",
        "harmony.chroma_fsharp",
        "harmony.chroma_g",
        "harmony.chroma_gsharp",
        "harmony.chroma_a",
        "harmony.chroma_asharp",
        "harmony.chroma_b",
        "harmony.key",
        "harmony.mode",
        "harmony.tonnetz_x",
        "harmony.tonnetz_y",
        "harmony.tonnetz_z",
        "harmony.tonnetz_u",
        "harmony.tonnetz_v",
        "harmony.tonnetz_w",
        "harmony.tonal_stability",
        "harmony.pitch_class_entropy",
    }
    actual = {f.name for f in result.feature_set}
    assert expected == actual


# --- Pipeline integration tests ---


def test_pipeline_integration(tmp_path: Path) -> None:
    """Verify pipeline runs HarmonyAnalyzer via default registry."""
    registry = build_default_registry()
    analyzers = registry.list_analyzers()
    names = [a.name for a in analyzers]
    assert "harmony" in names

    pipeline = AnalysisPipeline(registry)
    samples = _sine_wave(NOTE_FREQS["A"], duration=1.0).reshape(-1, 1)
    context = _make_context(samples, tmp_path)
    results = pipeline.run(context)

    harmony_results = [r for r in results if r.analyzer_name == "harmony"]
    assert len(harmony_results) == 1
    assert harmony_results[0].success is True
    assert len(harmony_results[0].feature_set) == 22


def test_persistence_integration(tmp_path: Path) -> None:
    """Verify pipeline with session persists harmony features to DB."""
    session = init_test_db(tmp_path)
    try:
        project = ProjectRepository.get_default_project(session)
        track = TrackRepository.create_track(
            session,
            project_id=project.id,
            relative_path="test.wav",
            original_filename="test.wav",
            sha256="hash_harmony_test",
            file_size=1000,
        )

        context = _make_context(
            _sine_wave(NOTE_FREQS["A"], duration=1.0).reshape(-1, 1),
            tmp_path,
        )
        context.track.id = track.id

        registry = build_default_registry()
        pipeline = AnalysisPipeline(registry)
        results = pipeline.run(context, session=session)
        session.commit()

        harmony_results = [r for r in results if r.analyzer_name == "harmony"]
        assert harmony_results[0].success is True

        runs = AnalysisRepository.list_runs(session, track.id)
        harmony_runs = [r for r in runs if r.analyzer_name == "harmony"]
        assert len(harmony_runs) == 1
        assert harmony_runs[0].success is True

        features = FeatureRepository.get_track_features(session, track.id)
        harmony_features = [
            f for f in features if f.analyzer_run.analyzer_name == "harmony"
        ]
        assert len(harmony_features) == 22
    finally:
        session.close()
