"""Music DNA similarity engine.

This module defines :class:`SimilarityEngine`, which compares two
:class:`~music_dna.vector.MusicDNAVector` objects and produces a
deterministic :class:`~similarity.result.SimilarityResult`.

The engine depends only on:
- MusicDNAVector
- MusicDNALayout
- similarity metrics

It has no knowledge of MusicDNA, databases, repositories, filesystems,
CLI, recommendation engines, or machine learning.
"""

from __future__ import annotations

import numpy as np

from music_dna.layout import MusicDNALayout
from music_dna.vector import MusicDNAVector
from similarity.exceptions import (
    SchemaMismatchError,
    VectorValidationError,
)
from similarity.explanation import SimilarityExplanation
from similarity.metrics import Metric, get_metric
from similarity.result import SimilarityResult


class SimilarityEngine:
    """Compares two MusicDNAVector objects and produces a similarity result.

    The engine is stateless and deterministic.  Comparing identical
    vectors always returns similarity 1.0 and distance 0.0.
    """

    SCHEMA_VERSION = "1.0"

    def __init__(self, layout: MusicDNALayout | None = None) -> None:
        """Initialize the engine.

        Args:
            layout: Optional layout instance.  If ``None``, a default
                :class:`MusicDNALayout` is created.
        """
        self.layout = layout or MusicDNALayout()

    def compare(
        self,
        a: MusicDNAVector,
        b: MusicDNAVector,
        metric: str = "cosine",
    ) -> SimilarityResult:
        """Compare two MusicDNAVector objects.

        Args:
            a: First vector.
            b: Second vector.
            metric: Metric name (``"cosine"``, ``"euclidean"``,
                ``"manhattan"``).  Defaults to ``"cosine"``.

        Returns:
            An immutable :class:`SimilarityResult`.

        Raises:
            SchemaMismatchError: If vectors have different schema
                versions.
            VectorValidationError: If vectors contain NaN, Inf, or
                have mismatched dimensions.
            UnsupportedMetricError: If the metric name is unknown.
        """
        self._validate_vectors(a, b)

        metric_impl: Metric = get_metric(metric)
        distance = metric_impl.distance(a.values, b.values)
        similarity = metric_impl.similarity(a.values, b.values)

        per_category = self._compute_per_category(a.values, b.values, metric_impl)

        top_contributors = SimilarityExplanation.generate(a, b, self.layout, top_n=5)

        return SimilarityResult(
            track_a=a.track_id,
            track_b=b.track_id,
            metric=metric,
            similarity_score=similarity,
            distance=distance,
            confidence=1.0,
            per_category_scores=per_category,
            top_contributors=top_contributors,
            vector_dimension=a.dimension,
            schema_version=a.schema_version,
        )

    def _validate_vectors(self, a: MusicDNAVector, b: MusicDNAVector) -> None:
        """Validate that two vectors are compatible for comparison.

        Args:
            a: First vector.
            b: Second vector.

        Raises:
            SchemaMismatchError: If schema versions differ.
            VectorValidationError: If dimensions differ or vectors
                contain NaN/Inf.
        """
        if a.schema_version != b.schema_version:
            raise SchemaMismatchError(
                f"Schema mismatch: {a.schema_version} vs {b.schema_version}"
            )
        if a.dimension != b.dimension:
            raise VectorValidationError(
                f"Dimension mismatch: {a.dimension} vs {b.dimension}"
            )
        if np.any(np.isnan(a.values)) or np.any(np.isnan(b.values)):
            raise VectorValidationError("Vectors must not contain NaN")
        if np.any(np.isinf(a.values)) or np.any(np.isinf(b.values)):
            raise VectorValidationError("Vectors must not contain Inf")

    def _compute_per_category(
        self,
        a: np.ndarray,
        b: np.ndarray,
        metric_impl: Metric,
    ) -> dict[str, float]:
        """Compute per-category similarity scores.

        Args:
            a: First vector array.
            b: Second vector array.
            metric_impl: Metric implementation to use.

        Returns:
            Dictionary mapping category name to similarity score.
        """
        categories = ["signal", "spectral", "dynamic", "rhythm", "harmony"]
        scores: dict[str, float] = {}

        for category in categories:
            indices = self.layout.indices(category)
            if not indices:
                continue
            sub_a = a[indices]
            sub_b = b[indices]
            scores[category] = metric_impl.similarity(
                sub_a.astype(np.float32),
                sub_b.astype(np.float32),
            )

        return scores
