"""Similarity metrics for Music DNA vector comparison.

This module defines pure mathematical metric functions and metric
classes that operate on ``numpy.ndarray`` objects.

All metrics validate input arrays for:
- Matching dimensions
- float32 dtype
- No NaN values
- No Inf values
- Non-empty vectors

Metric classes implement a common protocol with ``distance()`` and
``similarity()`` methods, allowing the engine to remain independent
of metric-specific normalization logic.
"""

from __future__ import annotations

from typing import Protocol

import numpy as np

from similarity.exceptions import VectorValidationError


def _validate_vectors(a: np.ndarray, b: np.ndarray) -> None:
    """Validate two vectors for comparison.

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

    Raises:
        VectorValidationError: If vectors are empty, have mismatched
            dimensions, contain NaN or Inf, or are not float32.
    """
    if a.size == 0 or b.size == 0:
        raise VectorValidationError("Vectors must not be empty")
    if a.shape != b.shape:
        raise VectorValidationError(f"Dimension mismatch: {a.shape} vs {b.shape}")
    if a.dtype != np.float32 or b.dtype != np.float32:
        raise VectorValidationError(
            f"Vectors must be float32, got {a.dtype} and {b.dtype}"
        )
    if np.any(np.isnan(a)) or np.any(np.isnan(b)):
        raise VectorValidationError("Vectors must not contain NaN")
    if np.any(np.isinf(a)) or np.any(np.isinf(b)):
        raise VectorValidationError("Vectors must not contain Inf")


def cosine_similarity(a: np.ndarray, b: np.ndarray) -> float:
    """Compute cosine similarity between two vectors.

    Returns a value in [0.0, 1.0].  Negative values are clamped to 0.0.

    Args:
        a: First vector (float32).
        b: Second vector (float32).

    Returns:
        Cosine similarity in [0.0, 1.0].
    """
    _validate_vectors(a, b)
    dot = float(np.dot(a, b))
    norm_a = float(np.linalg.norm(a))
    norm_b = float(np.linalg.norm(b))
    if norm_a == 0.0 or norm_b == 0.0:
        return 0.0
    sim = dot / (norm_a * norm_b)
    return max(0.0, min(1.0, sim))


def euclidean_distance(a: np.ndarray, b: np.ndarray) -> float:
    """Compute Euclidean (L2) distance between two vectors.

    Args:
        a: First vector (float32).
        b: Second vector (float32).

    Returns:
        Euclidean distance (>= 0.0).
    """
    _validate_vectors(a, b)
    return float(np.linalg.norm(a - b))


def manhattan_distance(a: np.ndarray, b: np.ndarray) -> float:
    """Compute Manhattan (L1) distance between two vectors.

    Args:
        a: First vector (float32).
        b: Second vector (float32).

    Returns:
        Manhattan distance (>= 0.0).
    """
    _validate_vectors(a, b)
    return float(np.sum(np.abs(a - b)))


class Metric(Protocol):
    """Common protocol for similarity metrics."""

    def distance(self, a: np.ndarray, b: np.ndarray) -> float: ...
    def similarity(self, a: np.ndarray, b: np.ndarray) -> float: ...


class CosineMetric:
    """Cosine similarity metric.

    Similarity is clamped to [0.0, 1.0].
    Distance is ``1.0 - similarity``.
    """

    def distance(self, a: np.ndarray, b: np.ndarray) -> float:
        """Return cosine distance (1 - similarity)."""
        return 1.0 - self.similarity(a, b)

    def similarity(self, a: np.ndarray, b: np.ndarray) -> float:
        """Return cosine similarity in [0.0, 1.0]."""
        return cosine_similarity(a, b)


class EuclideanMetric:
    """Euclidean (L2) distance metric.

    Distance is the L2 norm of the difference.
    Similarity is ``1 / (1 + distance)``.
    """

    def distance(self, a: np.ndarray, b: np.ndarray) -> float:
        """Return Euclidean distance."""
        return euclidean_distance(a, b)

    def similarity(self, a: np.ndarray, b: np.ndarray) -> float:
        """Return normalized similarity in [0.0, 1.0]."""
        d = euclidean_distance(a, b)
        return 1.0 / (1.0 + d)


class ManhattanMetric:
    """Manhattan (L1) distance metric.

    Distance is the L1 norm of the difference.
    Similarity is ``1 / (1 + distance)``.
    """

    def distance(self, a: np.ndarray, b: np.ndarray) -> float:
        """Return Manhattan distance."""
        return manhattan_distance(a, b)

    def similarity(self, a: np.ndarray, b: np.ndarray) -> float:
        """Return normalized similarity in [0.0, 1.0]."""
        d = manhattan_distance(a, b)
        return 1.0 / (1.0 + d)


_METRICS: dict[str, type[Metric]] = {
    "cosine": CosineMetric,
    "euclidean": EuclideanMetric,
    "manhattan": ManhattanMetric,
}


def get_metric(name: str) -> Metric:
    """Return a metric instance by name.

    Args:
        name: Metric name (``"cosine"``, ``"euclidean"``, ``"manhattan"``).

    Returns:
        A :class:`Metric` instance.

    Raises:
        UnsupportedMetricError: If the metric name is not recognized.
    """
    from similarity.exceptions import UnsupportedMetricError

    cls = _METRICS.get(name)
    if cls is None:
        raise UnsupportedMetricError(
            f"Unknown metric: '{name}'. "
            f"Supported: {', '.join(sorted(_METRICS.keys()))}"
        )
    return cls()
