"""Lesson 4 lab: norms, distances, cosine similarity, and metric choice.

Run:
    python 04-norms-distances-similarity.py

Requires:
    numpy

The implementations expose the arithmetic and validate their contracts. In a
production hot path, prefer tested vectorized library operations and benchmark
on representative shapes, dtypes, and hardware.
"""

from __future__ import annotations

import math

import numpy as np


def _vector(values: np.ndarray, name: str) -> np.ndarray:
    """Return a finite one-dimensional float array or raise a clear error."""
    vector = np.asarray(values, dtype=float)
    if vector.ndim != 1:
        raise ValueError(f"{name} must be one-dimensional; received {vector.shape}")
    if not np.all(np.isfinite(vector)):
        raise ValueError(f"{name} must contain only finite values")
    return vector


def lp_norm(values: np.ndarray, p: float) -> float:
    """Compute a genuine finite-dimensional Lp norm for p >= 1."""
    vector = _vector(values, "values")
    if not math.isfinite(p) or p < 1:
        raise ValueError("p must be finite and at least 1 for an Lp norm")
    return float(np.sum(np.abs(vector) ** p) ** (1.0 / p))


def minkowski_distance(left: np.ndarray, right: np.ndarray, p: float) -> float:
    """Compute the Lp distance between equal-shaped vectors."""
    x = _vector(left, "left")
    y = _vector(right, "right")
    if x.shape != y.shape:
        raise ValueError(f"vector shapes must match; received {x.shape} and {y.shape}")
    return lp_norm(x - y, p)


def cosine_similarity(left: np.ndarray, right: np.ndarray) -> float:
    """Return cosine similarity; reject zero vectors where the angle is undefined."""
    x = _vector(left, "left")
    y = _vector(right, "right")
    if x.shape != y.shape:
        raise ValueError(f"vector shapes must match; received {x.shape} and {y.shape}")
    denominator = lp_norm(x, 2) * lp_norm(y, 2)
    if denominator == 0.0:
        raise ValueError("cosine similarity is undefined for a zero vector")
    similarity = float(np.dot(x, y) / denominator)
    # Floating-point roundoff can produce values just beyond the mathematical range.
    return float(np.clip(similarity, -1.0, 1.0))


def pairwise_distances(samples: np.ndarray, p: float = 2) -> np.ndarray:
    """Return an [N, N] distance matrix for samples shaped [N, D]."""
    matrix = np.asarray(samples, dtype=float)
    if matrix.ndim != 2 or not np.all(np.isfinite(matrix)):
        raise ValueError("samples must be a finite [N, D] matrix")
    if not math.isfinite(p) or p < 1:
        raise ValueError("p must be finite and at least 1 for an Lp distance")
    differences = matrix[:, None, :] - matrix[None, :, :]  # [N, N, D]
    return np.sum(np.abs(differences) ** p, axis=-1) ** (1.0 / p)


def standardize(samples: np.ndarray) -> np.ndarray:
    """Standardize each feature; reject constant features explicitly."""
    matrix = np.asarray(samples, dtype=float)
    if matrix.ndim != 2 or not np.all(np.isfinite(matrix)):
        raise ValueError("samples must be a finite [N, D] matrix")
    mean = matrix.mean(axis=0, keepdims=True)
    scale = matrix.std(axis=0, keepdims=True)
    if np.any(scale == 0.0):
        raise ValueError("cannot standardize a constant feature")
    return (matrix - mean) / scale


def main() -> None:
    vector = np.array([3.0, -4.0])
    assert lp_norm(vector, 1) == 7.0
    assert lp_norm(vector, 2) == 5.0
    assert np.isclose(lp_norm(vector, 3), np.linalg.norm(vector, ord=3))

    x = np.array([1.0, 2.0])
    y = np.array([4.0, 6.0])
    assert minkowski_distance(x, y, 1) == 7.0
    assert minkowski_distance(x, y, 2) == 5.0
    assert np.isclose(cosine_similarity(x, 10.0 * x), 1.0)
    assert np.isclose(cosine_similarity(x, -x), -1.0)

    samples = np.array([[0.0, 0.0], [3.0, 4.0], [3.0, 0.0]])
    distances = pairwise_distances(samples)
    assert distances.shape == (3, 3)
    assert np.allclose(distances, distances.T)
    assert np.allclose(np.diag(distances), 0.0)
    assert np.isclose(distances[0, 1], 5.0)

    # A retrieval example: dot product rewards magnitude; cosine compares direction.
    query = np.array([1.0, 0.0])
    candidates = np.array([[10.0, 1.0], [1.0, 0.0], [0.8, 0.1]])
    dot_scores = candidates @ query
    cosine_scores = np.array([cosine_similarity(query, row) for row in candidates])
    assert int(np.argmax(dot_scores)) == 0
    assert int(np.argmax(cosine_scores)) == 1

    # Feature scale can reverse a nearest-centroid assignment.
    point = np.array([18.0, 55_000.0])  # age in years, income in dollars
    centroids = np.array([[20.0, 40_000.0], [50.0, 56_000.0]])
    raw_distances = np.linalg.norm(centroids - point, axis=1)
    population = np.vstack([point, centroids])
    scaled_population = standardize(population)
    scaled_distances = np.linalg.norm(
        scaled_population[1:] - scaled_population[0], axis=1
    )
    assert int(np.argmin(raw_distances)) == 1
    assert int(np.argmin(scaled_distances)) == 0

    try:
        cosine_similarity(np.zeros(2), x)
    except ValueError as error:
        assert "zero vector" in str(error)
    else:
        raise AssertionError("zero-vector cosine must be rejected")

    print("L1/L2 norms of [3, -4]:", lp_norm(vector, 1), lp_norm(vector, 2))
    print("pairwise L2 distances:\n", distances)
    print("retrieval dot scores:", dot_scores)
    print("retrieval cosine scores:", cosine_scores)
    print("raw centroid distances:", raw_distances)
    print("standardized centroid distances:", scaled_distances)

    # Try it yourself:
    # 1. Compare p=1, p=2, and p=4 rankings for a query and five candidates.
    # 2. Add one extreme feature value and compare raw versus robust scaling.
    # 3. L2-normalize candidates and verify dot-product and cosine rankings match.


if __name__ == "__main__":
    main()
