"""Unit tests for the Music DNA Similarity Engine.

Tests cover:
- Metric functions: cosine, euclidean, manhattan
- Metric classes: distance/similarity protocol
- Validation: NaN, Inf, dimension mismatch, empty vectors, dtype
- MusicDNALayout bidirectional mapping: get_feature, contains_index,
  category, indices
- SimilarityEngine: compare, per-category, determinism, identical vectors
- SimilarityExplanation: top contributors, deterministic ordering
- SimilarityResult: immutability
- Schema mismatch, unsupported metric
"""

from __future__ import annotations

from types import MappingProxyType

import numpy as np
import pytest

from music_dna.layout import MusicDNALayout
from music_dna.vector import MusicDNAVector, MusicDNAVectorMetadata
from similarity.engine import SimilarityEngine
from similarity.exceptions import (
    SchemaMismatchError,
    SimilarityError,
    UnsupportedMetricError,
    VectorValidationError,
)
from similarity.explanation import SimilarityExplanation
from similarity.metrics import (
    CosineMetric,
    EuclideanMetric,
    ManhattanMetric,
    cosine_similarity,
    euclidean_distance,
    get_metric,
    manhattan_distance,
)
from similarity.result import SimilarityResult

# --- Helpers ---


def _make_vector(
    track_id: str = "track-1",
    fill: float = 0.5,
    schema_version: str = "1.0",
) -> MusicDNAVector:
    """Create a synthetic MusicDNAVector filled with a constant value."""
    arr = np.full(512, fill, 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,
    )
    return MusicDNAVector(
        track_id=track_id,
        schema_version=schema_version,
        dimension=512,
        values=arr,
        metadata=metadata,
    )


def _make_vector_from_values(
    values: dict[int, float],
    track_id: str = "track-1",
) -> MusicDNAVector:
    """Create a MusicDNAVector with specific values at specific indices."""
    arr = np.zeros(512, dtype=np.float32)
    for idx, val in values.items():
        arr[idx] = val
    metadata = MusicDNAVectorMetadata(
        encoder_version="1.0.0",
        layout_version="1.0.0",
        created_at="2025-01-01T00:00:00Z",
        feature_count=len(values),
        reserved_dimensions=512 - len(values),
    )
    return MusicDNAVector(
        track_id=track_id,
        schema_version="1.0",
        dimension=512,
        values=arr,
        metadata=metadata,
    )


# --- Metric function tests ---


def test_cosine_similarity_identical_vectors() -> None:
    """Cosine similarity of identical vectors is 1.0."""
    a = np.ones(512, dtype=np.float32)
    b = np.ones(512, dtype=np.float32)
    assert cosine_similarity(a, b) == pytest.approx(1.0)


def test_cosine_similarity_orthogonal_vectors() -> None:
    """Cosine similarity of orthogonal vectors is 0.0."""
    a = np.zeros(512, dtype=np.float32)
    a[0] = 1.0
    b = np.zeros(512, dtype=np.float32)
    b[1] = 1.0
    assert cosine_similarity(a, b) == pytest.approx(0.0)


def test_cosine_similarity_different_vectors() -> None:
    """Cosine similarity of different vectors is between 0 and 1."""
    a = np.full(512, 0.5, dtype=np.float32)
    b = np.full(512, 1.0, dtype=np.float32)
    sim = cosine_similarity(a, b)
    assert 0.0 < sim <= 1.0


def test_cosine_similarity_clamped_negative() -> None:
    """Cosine similarity of opposite vectors is clamped to 0.0."""
    a = np.ones(512, dtype=np.float32)
    b = np.full(512, -1.0, dtype=np.float32)
    assert cosine_similarity(a, b) == 0.0


def test_cosine_similarity_zero_vector() -> None:
    """Cosine similarity with zero vector returns 0.0."""
    a = np.zeros(512, dtype=np.float32)
    b = np.ones(512, dtype=np.float32)
    assert cosine_similarity(a, b) == 0.0


def test_euclidean_distance_identical() -> None:
    """Euclidean distance of identical vectors is 0.0."""
    a = np.full(512, 0.5, dtype=np.float32)
    assert euclidean_distance(a, a) == pytest.approx(0.0)


def test_euclidean_distance_different() -> None:
    """Euclidean distance of different vectors is > 0."""
    a = np.zeros(512, dtype=np.float32)
    b = np.ones(512, dtype=np.float32)
    assert euclidean_distance(a, b) == pytest.approx(float(np.sqrt(512)))


def test_manhattan_distance_identical() -> None:
    """Manhattan distance of identical vectors is 0.0."""
    a = np.full(512, 0.5, dtype=np.float32)
    assert manhattan_distance(a, a) == pytest.approx(0.0)


def test_manhattan_distance_different() -> None:
    """Manhattan distance of different vectors is > 0."""
    a = np.zeros(512, dtype=np.float32)
    b = np.ones(512, dtype=np.float32)
    assert manhattan_distance(a, b) == pytest.approx(512.0)


# --- Metric validation tests ---


def test_cosine_nan_raises() -> None:
    """NaN in vectors raises VectorValidationError."""
    a = np.full(512, 0.5, dtype=np.float32)
    b = np.full(512, 0.5, dtype=np.float32)
    b[0] = np.nan
    with pytest.raises(VectorValidationError, match="NaN"):
        cosine_similarity(a, b)


def test_cosine_inf_raises() -> None:
    """Inf in vectors raises VectorValidationError."""
    a = np.full(512, 0.5, dtype=np.float32)
    b = np.full(512, 0.5, dtype=np.float32)
    b[0] = np.inf
    with pytest.raises(VectorValidationError, match="Inf"):
        cosine_similarity(a, b)


def test_cosine_dimension_mismatch_raises() -> None:
    """Dimension mismatch raises VectorValidationError."""
    a = np.ones(512, dtype=np.float32)
    b = np.ones(256, dtype=np.float32)
    with pytest.raises(VectorValidationError, match="Dimension"):
        cosine_similarity(a, b)


def test_cosine_empty_raises() -> None:
    """Empty vectors raise VectorValidationError."""
    a = np.array([], dtype=np.float32)
    b = np.array([], dtype=np.float32)
    with pytest.raises(VectorValidationError, match="empty"):
        cosine_similarity(a, b)


def test_euclidean_nan_raises() -> None:
    """NaN in euclidean raises VectorValidationError."""
    a = np.full(512, 0.5, dtype=np.float32)
    b = np.full(512, 0.5, dtype=np.float32)
    b[0] = np.nan
    with pytest.raises(VectorValidationError, match="NaN"):
        euclidean_distance(a, b)


def test_manhattan_inf_raises() -> None:
    """Inf in manhattan raises VectorValidationError."""
    a = np.full(512, 0.5, dtype=np.float32)
    b = np.full(512, 0.5, dtype=np.float32)
    b[0] = np.inf
    with pytest.raises(VectorValidationError, match="Inf"):
        manhattan_distance(a, b)


# --- Metric class tests ---


def test_cosine_metric_distance_and_similarity() -> None:
    """CosineMetric returns correct distance and similarity."""
    metric = CosineMetric()
    a = np.ones(512, dtype=np.float32)
    b = np.ones(512, dtype=np.float32)
    assert metric.similarity(a, b) == pytest.approx(1.0)
    assert metric.distance(a, b) == pytest.approx(0.0)


def test_euclidean_metric_distance_and_similarity() -> None:
    """EuclideanMetric returns correct distance and similarity."""
    metric = EuclideanMetric()
    a = np.zeros(512, dtype=np.float32)
    b = np.zeros(512, dtype=np.float32)
    assert metric.distance(a, b) == pytest.approx(0.0)
    assert metric.similarity(a, b) == pytest.approx(1.0)


def test_manhattan_metric_distance_and_similarity() -> None:
    """ManhattanMetric returns correct distance and similarity."""
    metric = ManhattanMetric()
    a = np.zeros(512, dtype=np.float32)
    b = np.zeros(512, dtype=np.float32)
    assert metric.distance(a, b) == pytest.approx(0.0)
    assert metric.similarity(a, b) == pytest.approx(1.0)


def test_get_metric_cosine() -> None:
    """get_metric returns CosineMetric for 'cosine'."""
    metric = get_metric("cosine")
    assert isinstance(metric, CosineMetric)


def test_get_metric_euclidean() -> None:
    """get_metric returns EuclideanMetric for 'euclidean'."""
    metric = get_metric("euclidean")
    assert isinstance(metric, EuclideanMetric)


def test_get_metric_manhattan() -> None:
    """get_metric returns ManhattanMetric for 'manhattan'."""
    metric = get_metric("manhattan")
    assert isinstance(metric, ManhattanMetric)


def test_get_metric_unknown_raises() -> None:
    """get_metric raises UnsupportedMetricError for unknown metric."""
    with pytest.raises(UnsupportedMetricError, match="unknown"):
        get_metric("unknown")


# --- Layout bidirectional mapping tests ---


def test_layout_get_feature() -> None:
    """Verify get_feature returns the correct identifier."""
    layout = MusicDNALayout()
    assert layout.get_feature(0) == "signal.rms"
    assert layout.get_feature(1) == "signal.peak"
    assert layout.get_feature(32) == "spectral.centroid"
    assert layout.get_feature(160) == "rhythm.tempo"
    assert layout.get_feature(245) == "harmony.pitch_class_entropy"


def test_layout_get_feature_not_found() -> None:
    """Verify KeyError for unmapped index."""
    layout = MusicDNALayout()
    with pytest.raises(KeyError, match="13"):
        layout.get_feature(13)


def test_layout_contains_index() -> None:
    """Verify contains_index returns True/False correctly."""
    layout = MusicDNALayout()
    assert layout.contains_index(0) is True
    assert layout.contains_index(13) is False
    assert layout.contains_index(511) is False


def test_layout_category() -> None:
    """Verify category returns correct category name."""
    layout = MusicDNALayout()
    assert layout.category(0) == "signal"
    assert layout.category(31) == "signal"
    assert layout.category(32) == "spectral"
    assert layout.category(96) == "dynamic"
    assert layout.category(160) == "rhythm"
    assert layout.category(224) == "harmony"
    assert layout.category(288) == "timbre"
    assert layout.category(480) == "reserved"


def test_layout_indices_for_signal() -> None:
    """Verify indices returns all signal indices."""
    layout = MusicDNALayout()
    signal_indices = layout.indices("signal")
    assert len(signal_indices) == 13
    assert 0 in signal_indices
    assert 12 in signal_indices
    assert all(0 <= i <= 31 for i in signal_indices)


def test_layout_indices_for_harmony() -> None:
    """Verify indices returns all harmony indices."""
    layout = MusicDNALayout()
    harmony_indices = layout.indices("harmony")
    assert len(harmony_indices) == 22
    assert 224 in harmony_indices
    assert 245 in harmony_indices
    assert all(224 <= i <= 287 for i in harmony_indices)


def test_layout_indices_for_empty_category() -> None:
    """Verify indices returns empty list for unmapped category."""
    layout = MusicDNALayout()
    assert layout.indices("timbre") == []


# --- SimilarityEngine tests ---


def test_engine_compare_identical_vectors() -> None:
    """Identical vectors produce similarity 1.0, distance 0.0."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=0.5)

    result = engine.compare(a, b, metric="cosine")

    assert result.similarity_score == pytest.approx(1.0)
    assert result.distance == pytest.approx(0.0)
    assert result.confidence == 1.0
    assert result.track_a == "track-1"
    assert result.track_b == "track-2"
    assert result.metric == "cosine"
    assert result.vector_dimension == 512


def test_engine_compare_different_vectors() -> None:
    """Different vectors produce similarity < 1.0."""
    engine = SimilarityEngine()
    a = _make_vector_from_values({0: 1.0, 1: 0.0}, "track-1")
    b = _make_vector_from_values({0: 0.0, 1: 1.0}, "track-2")

    result = engine.compare(a, b, metric="cosine")

    assert 0.0 <= result.similarity_score < 1.0
    assert result.distance > 0.0


def test_engine_compare_euclidean_identical() -> None:
    """Euclidean metric on identical vectors gives similarity 1.0."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=0.5)

    result = engine.compare(a, b, metric="euclidean")

    assert result.similarity_score == pytest.approx(1.0)
    assert result.distance == pytest.approx(0.0)


def test_engine_compare_manhattan_identical() -> None:
    """Manhattan metric on identical vectors gives similarity 1.0."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=0.5)

    result = engine.compare(a, b, metric="manhattan")

    assert result.similarity_score == pytest.approx(1.0)
    assert result.distance == pytest.approx(0.0)


def test_engine_per_category_scores() -> None:
    """Verify per-category scores are computed for all 5 categories."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=0.5)

    result = engine.compare(a, b, metric="cosine")

    assert "signal" in result.per_category_scores
    assert "spectral" in result.per_category_scores
    assert "dynamic" in result.per_category_scores
    assert "rhythm" in result.per_category_scores
    assert "harmony" in result.per_category_scores
    for score in result.per_category_scores.values():
        assert score == pytest.approx(1.0)


def test_engine_per_category_different_vectors() -> None:
    """Verify per-category scores differ for different vectors."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=1.0)

    result = engine.compare(a, b, metric="cosine")

    for score in result.per_category_scores.values():
        assert 0.0 <= score <= 1.0


def test_engine_deterministic() -> None:
    """Verify comparing same vectors twice produces identical results."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=0.7)

    result1 = engine.compare(a, b, metric="cosine")
    result2 = engine.compare(a, b, metric="cosine")

    assert result1.similarity_score == result2.similarity_score
    assert result1.distance == result2.distance
    assert result1.top_contributors == result2.top_contributors
    assert dict(result1.per_category_scores) == dict(result2.per_category_scores)


def test_engine_schema_mismatch_raises() -> None:
    """Verify SchemaMismatchError for different schema versions."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", schema_version="1.0")
    b = _make_vector("track-2", schema_version="2.0")

    with pytest.raises(SchemaMismatchError, match="Schema"):
        engine.compare(a, b)


def test_engine_nan_raises() -> None:
    """Verify VectorValidationError for NaN in vectors."""
    engine = SimilarityEngine()
    a = _make_vector("track-1")
    b = _make_vector("track-2")
    b.values.flags.writeable = True
    b.values[0] = np.nan
    b.values.flags.writeable = False

    with pytest.raises(VectorValidationError, match="NaN"):
        engine.compare(a, b)


def test_engine_inf_raises() -> None:
    """Verify VectorValidationError for Inf in vectors."""
    engine = SimilarityEngine()
    a = _make_vector("track-1")
    b = _make_vector("track-2")
    b.values.flags.writeable = True
    b.values[0] = np.inf
    b.values.flags.writeable = False

    with pytest.raises(VectorValidationError, match="Inf"):
        engine.compare(a, b)


def test_engine_unsupported_metric_raises() -> None:
    """Verify UnsupportedMetricError for unknown metric."""
    engine = SimilarityEngine()
    a = _make_vector("track-1")
    b = _make_vector("track-2")

    with pytest.raises(UnsupportedMetricError):
        engine.compare(a, b, metric="unknown")


def test_engine_default_layout() -> None:
    """Verify engine creates a default layout."""
    engine = SimilarityEngine()
    assert engine.layout is not None
    assert engine.layout.dimension() == 512


def test_engine_top_contributors() -> None:
    """Verify top_contributors is a list of feature identifiers."""
    engine = SimilarityEngine()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=1.0)

    result = engine.compare(a, b, metric="cosine")

    assert isinstance(result.top_contributors, list)
    assert len(result.top_contributors) == 5
    layout = MusicDNALayout()
    for contributor in result.top_contributors:
        assert layout.contains(contributor)


# --- SimilarityExplanation tests ---


def test_explanation_generates_ordered_list() -> None:
    """Verify explanation returns ordered list by largest difference."""
    layout = MusicDNALayout()
    a = _make_vector_from_values({0: 0.0, 1: 0.0, 2: 0.0}, "track-1")
    b = _make_vector_from_values({0: 0.1, 1: 0.5, 2: 0.3}, "track-2")

    explanation = SimilarityExplanation.generate(a, b, layout, top_n=3)

    assert len(explanation) == 3
    assert explanation[0] == "signal.peak"
    assert explanation[1] == "signal.dc_offset"
    assert explanation[2] == "signal.rms"


def test_explanation_deterministic() -> None:
    """Verify explanation is deterministic for same inputs."""
    layout = MusicDNALayout()
    a = _make_vector("track-1", fill=0.5)
    b = _make_vector("track-2", fill=1.0)

    result1 = SimilarityExplanation.generate(a, b, layout, top_n=5)
    result2 = SimilarityExplanation.generate(a, b, layout, top_n=5)

    assert result1 == result2


def test_explanation_ties_broken_alphabetically() -> None:
    """Verify ties in difference are broken by feature identifier."""
    layout = MusicDNALayout()
    a = _make_vector_from_values({0: 0.0, 1: 0.0}, "track-1")
    b = _make_vector_from_values({0: 0.5, 1: 0.5}, "track-2")

    explanation = SimilarityExplanation.generate(a, b, layout, top_n=2)

    assert explanation[0] == "signal.peak"
    assert explanation[1] == "signal.rms"


def test_explanation_top_n_limit() -> None:
    """Verify top_n limits the number of returned features."""
    layout = MusicDNALayout()
    a = _make_vector("track-1", fill=0.0)
    b = _make_vector("track-2", fill=1.0)

    explanation = SimilarityExplanation.generate(a, b, layout, top_n=3)

    assert len(explanation) == 3


# --- SimilarityResult tests ---


def test_result_is_immutable() -> None:
    """Verify SimilarityResult is frozen."""
    result = SimilarityResult(
        track_a="track-1",
        track_b="track-2",
        metric="cosine",
        similarity_score=0.95,
        distance=0.05,
        confidence=1.0,
        per_category_scores={"signal": 0.96},
        top_contributors=["signal.rms"],
        vector_dimension=512,
        schema_version="1.0",
    )
    with pytest.raises(AttributeError):
        result.track_a = "other"  # type: ignore[misc]


def test_result_per_category_scores_readonly() -> None:
    """Verify per_category_scores is a read-only mapping."""
    result = SimilarityResult(
        track_a="track-1",
        track_b="track-2",
        metric="cosine",
        similarity_score=0.95,
        distance=0.05,
        confidence=1.0,
        per_category_scores={"signal": 0.96},
        top_contributors=["signal.rms"],
    )
    assert isinstance(result.per_category_scores, MappingProxyType)
    with pytest.raises(TypeError):
        result.per_category_scores["new"] = 0.5  # type: ignore[index]


def test_all_exceptions_are_similarity_errors() -> None:
    """Verify all exceptions inherit from SimilarityError."""
    assert issubclass(VectorValidationError, SimilarityError)
    assert issubclass(SchemaMismatchError, SimilarityError)
    assert issubclass(UnsupportedMetricError, SimilarityError)
