"""Tests for the AudioDecoder module."""

from __future__ import annotations

from pathlib import Path

import numpy as np
import soundfile as sf

from core.exceptions import DecodingError
from decoder.decoded_audio import DecodedAudio
from decoder.decoder import AudioDecoder


def _generate_sine_wave(
    duration: float = 1.0,
    sample_rate: int = 22050,
    channels: int = 1,
    frequency: float = 440.0,
) -> np.ndarray:
    """Generate a deterministic sine wave signal.

    Args:
        duration: Duration in seconds.
        sample_rate: Sample rate in Hz.
        channels: Number of channels.
        frequency: Frequency of the sine wave in Hz.

    Returns:
        A float32 NumPy array with shape ``(n_frames, channels)``.
    """
    n_frames = int(sample_rate * duration)
    t = np.linspace(0, duration, n_frames, endpoint=False)
    mono = np.sin(2 * np.pi * frequency * t).astype(np.float32)
    if channels == 1:
        return mono.reshape(-1, 1)
    return np.column_stack([mono] * channels)


def _write_audio_file(
    path: Path,
    samples: np.ndarray,
    sample_rate: int,
    subtype: str = "PCM_16",
) -> None:
    """Write an audio file using soundfile.

    Args:
        path: Destination file path.
        samples: Audio samples as a NumPy array.
        sample_rate: Sample rate in Hz.
        subtype: Soundfile subtype string.
    """
    path.parent.mkdir(parents=True, exist_ok=True)
    sf.write(str(path), samples, sample_rate, subtype=subtype)


def test_wav_decoding(tmp_path: Path) -> None:
    """Verify that WAV files decode successfully."""
    samples = _generate_sine_wave(duration=1.0, sample_rate=22050, channels=1)
    file_path = tmp_path / "test.wav"
    _write_audio_file(file_path, samples, 22050, subtype="PCM_16")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert isinstance(decoded, DecodedAudio)
    assert decoded.sample_rate == 22050
    assert decoded.channels == 1
    assert decoded.duration == 1.0
    assert decoded.bit_depth == 16
    assert decoded.samples.dtype == np.float32
    assert decoded.samples.shape == (22050, 1)


def test_flac_decoding(tmp_path: Path) -> None:
    """Verify that FLAC files decode successfully."""
    samples = _generate_sine_wave(duration=0.5, sample_rate=44100, channels=1)
    file_path = tmp_path / "test.flac"
    _write_audio_file(file_path, samples, 44100, subtype="PCM_16")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert decoded.sample_rate == 44100
    assert decoded.channels == 1
    assert decoded.duration == 0.5
    assert decoded.bit_depth == 16


def test_aiff_decoding(tmp_path: Path) -> None:
    """Verify that AIFF files decode successfully."""
    samples = _generate_sine_wave(duration=0.5, sample_rate=44100, channels=1)
    file_path = tmp_path / "test.aiff"
    _write_audio_file(file_path, samples, 44100, subtype="PCM_16")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert decoded.sample_rate == 44100
    assert decoded.channels == 1
    assert decoded.duration == 0.5
    assert decoded.bit_depth == 16


def test_metadata_extraction(tmp_path: Path) -> None:
    """Verify that technical metadata is extracted correctly."""
    samples = _generate_sine_wave(duration=2.0, sample_rate=48000, channels=1)
    file_path = tmp_path / "meta.wav"
    _write_audio_file(file_path, samples, 48000, subtype="PCM_24")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert decoded.sample_rate == 48000
    assert decoded.channels == 1
    assert decoded.duration == 2.0
    assert decoded.bit_depth == 24


def test_stereo_preservation(tmp_path: Path) -> None:
    """Verify that stereo audio is not downmixed to mono."""
    samples = _generate_sine_wave(duration=1.0, sample_rate=22050, channels=2)
    file_path = tmp_path / "stereo.wav"
    _write_audio_file(file_path, samples, 22050, subtype="PCM_16")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert decoded.channels == 2
    assert decoded.samples.shape == (22050, 2)


def test_mono_preservation(tmp_path: Path) -> None:
    """Verify that mono audio is not expanded to stereo."""
    samples = _generate_sine_wave(duration=1.0, sample_rate=22050, channels=1)
    file_path = tmp_path / "mono.wav"
    _write_audio_file(file_path, samples, 22050, subtype="PCM_16")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert decoded.channels == 1
    assert decoded.samples.shape == (22050, 1)


def test_duration_calculation(tmp_path: Path) -> None:
    """Verify that duration is calculated correctly."""
    duration = 1.5
    sample_rate = 22050
    samples = _generate_sine_wave(
        duration=duration, sample_rate=sample_rate, channels=1
    )
    file_path = tmp_path / "dur.wav"
    _write_audio_file(file_path, samples, sample_rate, subtype="PCM_16")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert abs(decoded.duration - duration) < 0.01


def test_unsupported_file_handling(tmp_path: Path) -> None:
    """Verify that unsupported files raise DecodingError."""
    file_path = tmp_path / "test.txt"
    file_path.parent.mkdir(parents=True, exist_ok=True)
    file_path.write_bytes(b"not audio data")

    decoder = AudioDecoder()
    try:
        decoder.decode(file_path)
        raise AssertionError("Should have raised DecodingError")
    except DecodingError:
        pass


def test_corrupted_file_handling(tmp_path: Path) -> None:
    """Verify that corrupted audio files raise DecodingError."""
    file_path = tmp_path / "corrupt.wav"
    file_path.parent.mkdir(parents=True, exist_ok=True)
    file_path.write_bytes(b"RIFF\x00\x00\x00\x00corrupted data")

    decoder = AudioDecoder()
    try:
        decoder.decode(file_path)
        raise AssertionError("Should have raised DecodingError")
    except DecodingError:
        pass


def test_float32_normalization(tmp_path: Path) -> None:
    """Verify that samples are normalized to float32 range."""
    samples = _generate_sine_wave(duration=0.5, sample_rate=22050, channels=1)
    file_path = tmp_path / "norm.wav"
    _write_audio_file(file_path, samples, 22050, subtype="PCM_16")

    decoder = AudioDecoder()
    decoded = decoder.decode(file_path)

    assert decoded.samples.dtype == np.float32
    assert decoded.samples.min() >= -1.0
    assert decoded.samples.max() <= 1.0
