"""Unit tests for the Music DNA Vector Encoder.

Tests cover:
- MusicDNALayout: index lookup, contains, dimension, entries, reserved
- MusicDNAVector: immutability, float32 dtype, shape
- MusicDNAVectorMetadata: all fields stored
- MusicDNAEncoder: happy path, value conversion, reserved dims, determinism
- Validation: unknown features
- Full 55-feature encoding with correct index mapping
"""

from __future__ import annotations

from types import MappingProxyType

import numpy as np
import pytest

from music_dna.encoder import MusicDNAEncoder
from music_dna.exceptions import UnknownFeatureError
from music_dna.layout import MusicDNALayout
from music_dna.music_dna import MusicDNA, MusicDNAMetadata
from music_dna.vector import MusicDNAVector, MusicDNAVectorMetadata

# --- Helpers ---


def _make_dna(
    values: dict[str, float | str | int | bool],
    track_id: str = "track-1",
) -> MusicDNA:
    """Create a synthetic MusicDNA with the given values."""
    metadata = MusicDNAMetadata(
        builder_version="1.0.0",
        feature_count=len(values),
        normalization_version="1.0.0",
        source_analyzer_versions={"test": "1.0.0"},
    )
    return MusicDNA(
        track_id=track_id,
        schema_version="1.0",
        created_at="2025-01-01T00:00:00Z",
        values=MappingProxyType(values),
        metadata=metadata,
    )


def _make_all_values() -> dict[str, float | str | int | bool]:
    """Create values for all 55 registered features."""
    values: dict[str, float | str | int | bool] = {}
    layout = MusicDNALayout()
    for identifier, _index in layout.list_entries():
        if identifier in ("harmony.key", "harmony.mode"):
            values[identifier] = "A"
        elif identifier == "signal.channels":
            values[identifier] = 2
        elif identifier == "signal.sample_rate":
            values[identifier] = 44100
        else:
            values[identifier] = 0.5
    return values


# --- MusicDNALayout tests ---


def test_layout_dimension() -> None:
    """Verify layout dimension is 512."""
    layout = MusicDNALayout()
    assert layout.dimension() == 512


def test_layout_feature_count() -> None:
    """Verify layout has 55 mapped features."""
    layout = MusicDNALayout()
    assert layout.feature_count() == 55


def test_layout_get_index() -> None:
    """Verify specific index mappings."""
    layout = MusicDNALayout()
    assert layout.get_index("signal.rms") == 0
    assert layout.get_index("signal.peak") == 1
    assert layout.get_index("signal.dc_offset") == 2
    assert layout.get_index("spectral.centroid") == 32
    assert layout.get_index("spectral.bandwidth") == 33
    assert layout.get_index("dynamic.crest_factor") == 96
    assert layout.get_index("rhythm.tempo") == 160
    assert layout.get_index("harmony.chroma_c") == 224
    assert layout.get_index("harmony.chroma_b") == 235
    assert layout.get_index("harmony.tonnetz_x") == 236
    assert layout.get_index("harmony.tonnetz_w") == 241
    assert layout.get_index("harmony.key") == 242
    assert layout.get_index("harmony.pitch_class_entropy") == 245


def test_layout_get_index_not_found() -> None:
    """Verify KeyError for unknown identifier."""
    layout = MusicDNALayout()
    with pytest.raises(KeyError, match="unknown.feature"):
        layout.get_index("unknown.feature")


def test_layout_contains() -> None:
    """Verify contains returns True/False correctly."""
    layout = MusicDNALayout()
    assert layout.contains("signal.rms") is True
    assert layout.contains("unknown.feature") is False


def test_layout_list_entries_sorted_by_index() -> None:
    """Verify list_entries returns entries sorted by index."""
    layout = MusicDNALayout()
    entries = layout.list_entries()
    indices = [idx for _, idx in entries]
    assert indices == sorted(indices)
    assert len(entries) == 55


def test_layout_reserved_indices() -> None:
    """Verify reserved indices are correct."""
    layout = MusicDNALayout()
    reserved = layout.reserved_indices()
    assert len(reserved) == 512 - 55
    assert 13 in reserved
    assert 0 not in reserved
    assert 245 not in reserved
    assert 511 in reserved


def test_layout_version() -> None:
    """Verify layout has a version string."""
    assert MusicDNALayout.LAYOUT_VERSION == "1.0.0"


# --- MusicDNAVector tests ---


def test_vector_is_immutable() -> None:
    """Verify the values array is read-only."""
    arr = np.zeros(512, dtype=np.float32)
    metadata = MusicDNAVectorMetadata(
        encoder_version="1.0.0",
        layout_version="1.0.0",
        created_at="2025-01-01T00:00:00Z",
        feature_count=55,
        reserved_dimensions=457,
    )
    vec = MusicDNAVector(
        track_id="track-1",
        schema_version="1.0",
        dimension=512,
        values=arr,
        metadata=metadata,
    )
    assert vec.values.flags.writeable is False
    with pytest.raises(ValueError):
        vec.values[0] = 1.0


def test_vector_float32_dtype() -> None:
    """Verify the values array has float32 dtype."""
    arr = np.zeros(512, dtype=np.float32)
    metadata = MusicDNAVectorMetadata(
        encoder_version="1.0.0",
        layout_version="1.0.0",
        created_at="2025-01-01T00:00:00Z",
        feature_count=55,
        reserved_dimensions=457,
    )
    vec = MusicDNAVector(
        track_id="track-1",
        schema_version="1.0",
        dimension=512,
        values=arr,
        metadata=metadata,
    )
    assert vec.values.dtype == np.float32


def test_vector_shape() -> None:
    """Verify the values array has shape (512,)."""
    arr = np.zeros(512, dtype=np.float32)
    metadata = MusicDNAVectorMetadata(
        encoder_version="1.0.0",
        layout_version="1.0.0",
        created_at="2025-01-01T00:00:00Z",
        feature_count=55,
        reserved_dimensions=457,
    )
    vec = MusicDNAVector(
        track_id="track-1",
        schema_version="1.0",
        dimension=512,
        values=arr,
        metadata=metadata,
    )
    assert vec.values.shape == (512,)


def test_vector_metadata_stored() -> None:
    """Verify all metadata fields are stored."""
    arr = np.zeros(512, dtype=np.float32)
    metadata = MusicDNAVectorMetadata(
        encoder_version="1.0.0",
        layout_version="1.0.0",
        created_at="2025-01-01T00:00:00Z",
        feature_count=55,
        reserved_dimensions=457,
    )
    vec = MusicDNAVector(
        track_id="track-1",
        schema_version="1.0",
        dimension=512,
        values=arr,
        metadata=metadata,
    )
    assert vec.metadata.encoder_version == "1.0.0"
    assert vec.metadata.layout_version == "1.0.0"
    assert vec.metadata.created_at == "2025-01-01T00:00:00Z"
    assert vec.metadata.feature_count == 55
    assert vec.metadata.reserved_dimensions == 457


# --- MusicDNAEncoder tests ---


def test_encoder_happy_path() -> None:
    """Verify encoder produces a valid MusicDNAVector from MusicDNA."""
    values = _make_all_values()
    dna = _make_dna(values)
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    assert vec.track_id == "track-1"
    assert vec.schema_version == "1.0"
    assert vec.dimension == 512
    assert vec.values.dtype == np.float32
    assert vec.values.shape == (512,)


def test_encoder_float_values_at_correct_indices() -> None:
    """Verify float values are stored at the correct indices."""
    values = _make_all_values()
    dna = _make_dna(values)
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    layout = MusicDNALayout()
    for identifier, index in layout.list_entries():
        if identifier in ("harmony.key", "harmony.mode"):
            assert vec.values[index] == 0.0
        elif identifier == "signal.channels":
            assert vec.values[index] == 2.0
        elif identifier == "signal.sample_rate":
            assert vec.values[index] == 44100.0
        else:
            assert vec.values[index] == 0.5


def test_encoder_bool_values() -> None:
    """Verify boolean values are converted to 0.0/1.0."""
    dna = _make_dna({"signal.rms": True, "signal.peak": False})
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    assert vec.values[0] == 1.0
    assert vec.values[1] == 0.0


def test_encoder_string_values_are_zero() -> None:
    """Verify categorical (string) values are stored as 0.0."""
    dna = _make_dna({"harmony.key": "A", "harmony.mode": "minor"})
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    assert vec.values[242] == 0.0
    assert vec.values[243] == 0.0


def test_encoder_int_values() -> None:
    """Verify int values are stored as floats."""
    dna = _make_dna({"signal.sample_rate": 48000, "signal.channels": 2})
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    assert vec.values[7] == 48000.0
    assert vec.values[8] == 2.0


def test_encoder_reserved_dimensions_are_zero() -> None:
    """Verify all reserved dimensions are 0.0."""
    values = _make_all_values()
    dna = _make_dna(values)
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    layout = MusicDNALayout()
    for idx in layout.reserved_indices():
        assert vec.values[idx] == 0.0


def test_encoder_reserved_dimensions_no_nan() -> None:
    """Verify no NaN values in the vector."""
    values = _make_all_values()
    dna = _make_dna(values)
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    assert not np.any(np.isnan(vec.values))


def test_encoder_deterministic() -> None:
    """Verify identical MusicDNA produces byte-identical vectors."""
    values = _make_all_values()
    dna1 = _make_dna(values)
    dna2 = _make_dna(values)
    encoder = MusicDNAEncoder()

    vec1 = encoder.encode(dna1, "2025-01-01T00:00:00Z")
    vec2 = encoder.encode(dna2, "2025-01-01T00:00:00Z")

    assert np.array_equal(vec1.values, vec2.values)
    assert vec1.values.tobytes() == vec2.values.tobytes()


def test_encoder_unknown_feature_raises() -> None:
    """Verify UnknownFeatureError for features not in the layout."""
    dna = _make_dna({"unknown.feature": 0.5})
    encoder = MusicDNAEncoder()

    with pytest.raises(UnknownFeatureError, match="unknown.feature"):
        encoder.encode(dna, "2025-01-01T00:00:00Z")


def test_encoder_metadata() -> None:
    """Verify encoder metadata is correct."""
    values = _make_all_values()
    dna = _make_dna(values)
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    assert vec.metadata.encoder_version == "1.0.0"
    assert vec.metadata.layout_version == "1.0.0"
    assert vec.metadata.created_at == "2025-01-01T00:00:00Z"
    assert vec.metadata.feature_count == 55
    assert vec.metadata.reserved_dimensions == 512 - 55


def test_encoder_default_layout() -> None:
    """Verify encoder creates a default layout when none is provided."""
    encoder = MusicDNAEncoder()
    assert encoder.layout is not None
    assert encoder.layout.dimension() == 512


def test_encoder_all_55_features_mapped() -> None:
    """Verify all 55 features are written to the vector."""
    values = _make_all_values()
    dna = _make_dna(values)
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    layout = MusicDNALayout()
    non_zero = 0
    for identifier, index in layout.list_entries():
        if identifier not in ("harmony.key", "harmony.mode") and (
            vec.values[index] != 0.0
        ):
            non_zero += 1
    assert non_zero == 53  # 55 total - 2 categorical strings


def test_encoder_vector_always_512() -> None:
    """Verify output vector is always 512-dimensional."""
    dna = _make_dna({"signal.rms": 0.5})
    encoder = MusicDNAEncoder()

    vec = encoder.encode(dna, "2025-01-01T00:00:00Z")

    assert vec.dimension == 512
    assert vec.values.shape == (512,)


def test_encoder_version() -> None:
    """Verify encoder has a version string."""
    assert MusicDNAEncoder.ENCODER_VERSION == "1.0.0"
