"""Inspectable sparse, dense, approximate, and compressed vector indexes.

These implementations favor auditability over throughput.  They are useful as
correctness baselines, teaching implementations, and fixtures for comparing
production engines.  In particular, the IVF index exposes its centroids and
assignments, while quantizers report reconstruction error and true storage
components instead of hiding compression behind a vendor API.
"""

import math
from collections import Counter, defaultdict
from dataclasses import dataclass
from typing import (
    AbstractSet,
    DefaultDict,
    Dict,
    List,
    Mapping,
    Optional,
    Protocol,
    Sequence,
    Tuple,
    Union,
)

from .models import Chunk, SearchResult
from .text import tokenize


Vector = Tuple[float, ...]
VectorInput = Union[Sequence[Sequence[float]], Mapping[str, Sequence[float]]]


def _normalize(vector: Sequence[float]) -> Vector:
    norm = math.sqrt(sum(value * value for value in vector))
    if norm == 0.0:
        return tuple(0.0 for _ in vector)
    return tuple(value / norm for value in vector)


def _cosine(left: Sequence[float], right: Sequence[float]) -> float:
    return sum(a * b for a, b in zip(left, right))


def _squared_distance(left: Sequence[float], right: Sequence[float]) -> float:
    return sum((a - b) ** 2 for a, b in zip(left, right))


def _materialize_vectors(chunks: Sequence[Chunk], vectors: VectorInput) -> Tuple[Vector, ...]:
    if isinstance(vectors, Mapping):
        try:
            materialized = tuple(
                tuple(float(value) for value in vectors[chunk.id]) for chunk in chunks
            )
        except KeyError as error:
            raise ValueError(f"missing vector for chunk {error.args[0]}") from error
    else:
        materialized = tuple(tuple(float(value) for value in vector) for vector in vectors)
    if len(materialized) != len(chunks):
        raise ValueError("chunks and vectors must have equal length")
    dimensions = len(materialized[0]) if materialized else 0
    if materialized and dimensions == 0:
        raise ValueError("vectors must have at least one dimension")
    if any(len(vector) != dimensions for vector in materialized):
        raise ValueError("all vectors must have equal dimensionality")
    if any(not math.isfinite(value) for vector in materialized for value in vector):
        raise ValueError("vectors must contain only finite values")
    return materialized


@dataclass(frozen=True)
class Posting:
    """One compressed-index concept: document ordinal and term frequency."""

    document_ordinal: int
    term_frequency: int


class InvertedIndex:
    """A real postings-list BM25 index.

    Query evaluation touches only postings for query terms, unlike the package's
    intentionally simple scan-based BM25 retriever.  Robertson/Sparck Jones IDF
    is kept positive with ``log(1 + ratio)``.
    """

    name = "inverted-bm25"

    def __init__(self, chunks: Sequence[Chunk], k1: float = 1.2, b: float = 0.75) -> None:
        if k1 <= 0.0:
            raise ValueError("k1 must be positive")
        if not 0.0 <= b <= 1.0:
            raise ValueError("b must be between zero and one")
        self.chunks = tuple(chunks)
        self.k1 = k1
        self.b = b
        postings: DefaultDict[str, List[Posting]] = defaultdict(list)
        lengths: List[int] = []
        for ordinal, chunk in enumerate(self.chunks):
            frequencies = Counter(tokenize(chunk.title + " " + chunk.text))
            lengths.append(sum(frequencies.values()))
            for term, frequency in frequencies.items():
                postings[term].append(Posting(ordinal, frequency))
        self.postings: Mapping[str, Tuple[Posting, ...]] = {
            term: tuple(items) for term, items in postings.items()
        }
        self.document_lengths = tuple(lengths)
        self.average_document_length = sum(lengths) / len(lengths) if lengths else 0.0

    def idf(self, term: str) -> float:
        """Return positive RSJ inverse document frequency for a term."""

        frequency = len(self.postings.get(term, ()))
        total = len(self.chunks)
        return math.log(1.0 + (total - frequency + 0.5) / (frequency + 0.5))

    def search(
        self,
        query: str,
        k: int = 5,
        allowed_chunk_ids: Optional[AbstractSet[str]] = None,
    ) -> List[SearchResult]:
        """Accumulate BM25 scores from query-term postings and rank deterministically."""

        if k <= 0 or not self.average_document_length:
            return []
        scores: DefaultDict[int, float] = defaultdict(float)
        query_frequencies = Counter(tokenize(query))
        for term, query_frequency in query_frequencies.items():
            idf = self.idf(term)
            for posting in self.postings.get(term, ()):
                chunk = self.chunks[posting.document_ordinal]
                if allowed_chunk_ids is not None and chunk.id not in allowed_chunk_ids:
                    continue
                length_ratio = self.document_lengths[posting.document_ordinal]
                length_ratio /= self.average_document_length
                denominator = posting.term_frequency + self.k1 * (
                    1.0 - self.b + self.b * length_ratio
                )
                saturation = posting.term_frequency * (self.k1 + 1.0) / denominator
                scores[posting.document_ordinal] += (
                    idf * saturation * (1.0 + math.log(query_frequency))
                )
        ranked = sorted(scores.items(), key=lambda item: (-item[1], self.chunks[item[0]].id))
        return [
            SearchResult(
                chunk=self.chunks[ordinal],
                score=score,
                rank=rank,
                retriever=self.name,
                component_scores={self.name: score},
            )
            for rank, (ordinal, score) in enumerate(ranked[:k], start=1)
            if score > 0.0
        ]


class ExactCosineIndex:
    """Brute-force cosine index used as an exact ANN recall baseline."""

    name = "exact-cosine"

    def __init__(self, chunks: Sequence[Chunk], vectors: VectorInput) -> None:
        self.chunks = tuple(chunks)
        materialized = _materialize_vectors(self.chunks, vectors)
        self.vectors = tuple(_normalize(vector) for vector in materialized)
        self.dimensions = len(self.vectors[0]) if self.vectors else 0

    def search(
        self,
        query_vector: Sequence[float],
        k: int = 5,
        allowed_chunk_ids: Optional[AbstractSet[str]] = None,
    ) -> List[SearchResult]:
        """Return exact cosine neighbors, applying filters before top-k selection."""

        if k <= 0 or not self.chunks:
            return []
        if len(query_vector) != self.dimensions:
            raise ValueError("query dimensionality does not match the index")
        query = _normalize(query_vector)
        scored = [
            (index, _cosine(query, vector))
            for index, vector in enumerate(self.vectors)
            if allowed_chunk_ids is None or self.chunks[index].id in allowed_chunk_ids
        ]
        scored.sort(key=lambda item: (-item[1], self.chunks[item[0]].id))
        return [
            SearchResult(
                chunk=self.chunks[index],
                score=score,
                rank=rank,
                retriever=self.name,
                component_scores={self.name: score},
            )
            for rank, (index, score) in enumerate(scored[:k], start=1)
        ]


def _mean(vectors: Sequence[Sequence[float]], dimensions: int) -> Vector:
    if not vectors:
        return tuple(0.0 for _ in range(dimensions))
    return tuple(sum(vector[d] for vector in vectors) / len(vectors) for d in range(dimensions))


def _kmeans(
    vectors: Sequence[Vector],
    clusters: int,
    iterations: int,
    cosine: bool,
) -> Tuple[Tuple[Vector, ...], Tuple[int, ...]]:
    """Deterministic farthest-first initialization plus Lloyd iterations."""

    if not vectors:
        return (), ()
    cluster_count = min(clusters, len(vectors))
    centroids: List[Vector] = [vectors[0]]
    while len(centroids) < cluster_count:
        distances = [
            min(_squared_distance(vector, centroid) for centroid in centroids)
            for vector in vectors
        ]
        next_index = max(range(len(vectors)), key=lambda index: (distances[index], -index))
        centroids.append(vectors[next_index])
    assignments = [0] * len(vectors)
    for _ in range(iterations):
        for index, vector in enumerate(vectors):
            if cosine:
                assignments[index] = max(
                    range(cluster_count),
                    key=lambda cluster: (_cosine(vector, centroids[cluster]), -cluster),
                )
            else:
                assignments[index] = min(
                    range(cluster_count),
                    key=lambda cluster: (_squared_distance(vector, centroids[cluster]), cluster),
                )
        updated: List[Vector] = []
        for cluster in range(cluster_count):
            members = [
                vectors[index]
                for index, value in enumerate(assignments)
                if value == cluster
            ]
            centroid = _mean(members, len(vectors[0])) if members else centroids[cluster]
            updated.append(_normalize(centroid) if cosine else centroid)
        if updated == centroids:
            break
        centroids = updated
    return tuple(centroids), tuple(assignments)


class IVFCoarseIndex:
    """Inverted-file ANN index with transparent coarse k-means cells.

    Vectors and centroids are L2-normalized.  Search probes the ``nprobe`` most
    similar centroids, then computes exact cosine scores only for vectors in
    those cells.  Increasing ``nprobe`` trades latency for recall.
    """

    name = "ivf-cosine"

    def __init__(
        self,
        chunks: Sequence[Chunk],
        vectors: VectorInput,
        nlist: int = 8,
        iterations: int = 12,
    ) -> None:
        if nlist <= 0:
            raise ValueError("nlist must be positive")
        if iterations <= 0:
            raise ValueError("iterations must be positive")
        self.chunks = tuple(chunks)
        materialized = _materialize_vectors(self.chunks, vectors)
        self.vectors = tuple(_normalize(vector) for vector in materialized)
        self.dimensions = len(self.vectors[0]) if self.vectors else 0
        self.centroids, self.assignments = _kmeans(self.vectors, nlist, iterations, True)
        lists: DefaultDict[int, List[int]] = defaultdict(list)
        for ordinal, cluster in enumerate(self.assignments):
            lists[cluster].append(ordinal)
        self.inverted_lists: Mapping[int, Tuple[int, ...]] = {
            cluster: tuple(ordinals) for cluster, ordinals in lists.items()
        }
        self.nlist = len(self.centroids)

    def search(
        self,
        query_vector: Sequence[float],
        k: int = 5,
        nprobe: int = 1,
        allowed_chunk_ids: Optional[AbstractSet[str]] = None,
    ) -> List[SearchResult]:
        """Search selected coarse cells and expose cell/probe information."""

        if k <= 0 or not self.chunks:
            return []
        if len(query_vector) != self.dimensions:
            raise ValueError("query dimensionality does not match the index")
        if nprobe <= 0:
            raise ValueError("nprobe must be positive")
        query = _normalize(query_vector)
        centroid_scores = sorted(
            (
                (cluster, _cosine(query, centroid))
                for cluster, centroid in enumerate(self.centroids)
            ),
            key=lambda item: (-item[1], item[0]),
        )
        probed = centroid_scores[: min(nprobe, self.nlist)]
        coarse_scores = dict(probed)
        candidate_ordinals = [
            ordinal
            for cluster, _ in probed
            for ordinal in self.inverted_lists.get(cluster, ())
            if allowed_chunk_ids is None or self.chunks[ordinal].id in allowed_chunk_ids
        ]
        scored = [
            (ordinal, _cosine(query, self.vectors[ordinal]))
            for ordinal in candidate_ordinals
        ]
        scored.sort(key=lambda item: (-item[1], self.chunks[item[0]].id))
        return [
            SearchResult(
                chunk=self.chunks[ordinal],
                score=score,
                rank=rank,
                retriever=self.name,
                component_scores={
                    "coarse_score": coarse_scores[self.assignments[ordinal]],
                    "ivf_cell": float(self.assignments[ordinal]),
                    "nprobe": float(min(nprobe, self.nlist)),
                },
            )
            for rank, (ordinal, score) in enumerate(scored[:k], start=1)
        ]


class VectorQuantizer(Protocol):
    """Structural interface consumed by :func:`audit_quantization`."""

    dimensions: int

    @property
    def encoded_bytes_per_vector(self) -> int:
        """Return packed code bytes for one vector."""

        ...

    @property
    def codebook_bytes(self) -> int:
        """Return shared parameter storage in bytes."""

        ...

    def encode(self, vector: Sequence[float]) -> Tuple[int, ...]:
        """Encode one vector to integer codes."""

        ...

    def decode(self, codes: Sequence[int]) -> Vector:
        """Reconstruct one vector from integer codes."""

        ...


class ScalarQuantizer:
    """Per-dimension uniform min/max scalar quantization."""

    def __init__(self, bits: int = 8) -> None:
        if not 1 <= bits <= 16:
            raise ValueError("bits must be between 1 and 16")
        self.bits = bits
        self.levels = (1 << bits) - 1
        self.minimums: Vector = ()
        self.maximums: Vector = ()
        self.dimensions = 0

    def fit(self, vectors: Sequence[Sequence[float]]) -> "ScalarQuantizer":
        """Learn an independent numeric range for every dimension."""

        materialized = tuple(tuple(float(value) for value in vector) for vector in vectors)
        if not materialized or not materialized[0]:
            raise ValueError("at least one non-empty vector is required")
        dimensions = len(materialized[0])
        if any(len(vector) != dimensions for vector in materialized):
            raise ValueError("all vectors must have equal dimensionality")
        self.minimums = tuple(min(vector[d] for vector in materialized) for d in range(dimensions))
        self.maximums = tuple(max(vector[d] for vector in materialized) for d in range(dimensions))
        self.dimensions = dimensions
        return self

    def encode(self, vector: Sequence[float]) -> Tuple[int, ...]:
        """Quantize values into ``2**bits`` uniformly spaced levels."""

        if len(vector) != self.dimensions or not self.dimensions:
            raise ValueError("vector dimensionality does not match the fitted quantizer")
        codes: List[int] = []
        for value, minimum, maximum in zip(vector, self.minimums, self.maximums):
            if maximum == minimum:
                codes.append(0)
            else:
                scaled = (value - minimum) / (maximum - minimum)
                codes.append(max(0, min(self.levels, int(round(scaled * self.levels)))))
        return tuple(codes)

    def decode(self, codes: Sequence[int]) -> Vector:
        """Reconstruct the center represented by each scalar code."""

        if len(codes) != self.dimensions or not self.dimensions:
            raise ValueError("code dimensionality does not match the fitted quantizer")
        if any(code < 0 or code > self.levels for code in codes):
            raise ValueError("scalar code is outside the configured range")
        return tuple(
            minimum + (code / self.levels) * (maximum - minimum)
            for code, minimum, maximum in zip(codes, self.minimums, self.maximums)
        )

    @property
    def encoded_bytes_per_vector(self) -> int:
        return math.ceil(self.dimensions * self.bits / 8)

    @property
    def codebook_bytes(self) -> int:
        return self.dimensions * 2 * 4


class ProductQuantizer:
    """Product quantization with deterministic k-means sub-codebooks.

    Dimensions are split as evenly as possible among ``subquantizers``.  Each
    subvector is represented by the id of its nearest centroid.  This is a
    compact PQ-style reference implementation; optimized systems use SIMD ADC
    lookup tables and often learn an OPQ rotation first.
    """

    def __init__(self, subquantizers: int = 4, bits: int = 4, iterations: int = 12) -> None:
        if subquantizers <= 0:
            raise ValueError("subquantizers must be positive")
        if not 1 <= bits <= 8:
            raise ValueError("bits must be between 1 and 8")
        if iterations <= 0:
            raise ValueError("iterations must be positive")
        self.subquantizers = subquantizers
        self.bits = bits
        self.iterations = iterations
        self.dimensions = 0
        self.boundaries: Tuple[Tuple[int, int], ...] = ()
        self.codebooks: Tuple[Tuple[Vector, ...], ...] = ()

    def fit(self, vectors: Sequence[Sequence[float]]) -> "ProductQuantizer":
        """Train one Euclidean k-means codebook per dimension block."""

        materialized = tuple(tuple(float(value) for value in vector) for vector in vectors)
        if not materialized or not materialized[0]:
            raise ValueError("at least one non-empty vector is required")
        dimensions = len(materialized[0])
        if any(len(vector) != dimensions for vector in materialized):
            raise ValueError("all vectors must have equal dimensionality")
        if self.subquantizers > dimensions:
            raise ValueError("subquantizers must not exceed vector dimensions")
        boundaries: List[Tuple[int, int]] = []
        for block in range(self.subquantizers):
            start = block * dimensions // self.subquantizers
            end = (block + 1) * dimensions // self.subquantizers
            boundaries.append((start, end))
        codebooks: List[Tuple[Vector, ...]] = []
        cluster_count = min(1 << self.bits, len(materialized))
        for start, end in boundaries:
            subvectors = tuple(vector[start:end] for vector in materialized)
            centroids, _ = _kmeans(subvectors, cluster_count, self.iterations, False)
            codebooks.append(centroids)
        self.dimensions = dimensions
        self.boundaries = tuple(boundaries)
        self.codebooks = tuple(codebooks)
        return self

    def encode(self, vector: Sequence[float]) -> Tuple[int, ...]:
        """Return the nearest centroid id for each subvector."""

        if len(vector) != self.dimensions or not self.codebooks:
            raise ValueError("vector dimensionality does not match the fitted quantizer")
        codes: List[int] = []
        for (start, end), codebook in zip(self.boundaries, self.codebooks):
            subvector = vector[start:end]
            code = min(
                range(len(codebook)),
                key=lambda index: (_squared_distance(subvector, codebook[index]), index),
            )
            codes.append(code)
        return tuple(codes)

    def decode(self, codes: Sequence[int]) -> Vector:
        """Concatenate selected centroids into a reconstructed vector."""

        if len(codes) != len(self.codebooks) or not self.codebooks:
            raise ValueError("code count does not match the fitted quantizer")
        output: List[float] = []
        for code, codebook in zip(codes, self.codebooks):
            if code < 0 or code >= len(codebook):
                raise ValueError("product-quantizer code is outside its codebook")
            output.extend(codebook[code])
        return tuple(output)

    @property
    def encoded_bytes_per_vector(self) -> int:
        return math.ceil(self.subquantizers * self.bits / 8)

    @property
    def codebook_bytes(self) -> int:
        return sum(len(codebook) * len(codebook[0]) * 4 for codebook in self.codebooks if codebook)


@dataclass(frozen=True)
class QuantizationAudit:
    """Accuracy and storage report for reconstructed vectors."""

    vector_count: int
    dimensions: int
    original_bytes: int
    encoded_bytes: int
    codebook_bytes: int
    compression_ratio: float
    mean_squared_error: float
    mean_cosine_similarity: float


def audit_quantization(
    vectors: Sequence[Sequence[float]],
    quantizer: VectorQuantizer,
    float_bytes: int = 4,
) -> QuantizationAudit:
    """Measure reconstruction quality and total model-plus-code storage."""

    if float_bytes <= 0:
        raise ValueError("float_bytes must be positive")
    materialized = tuple(tuple(float(value) for value in vector) for vector in vectors)
    if not materialized:
        raise ValueError("at least one vector is required")
    if any(len(vector) != quantizer.dimensions for vector in materialized):
        raise ValueError("vector dimensionality does not match the quantizer")
    reconstructed = tuple(quantizer.decode(quantizer.encode(vector)) for vector in materialized)
    squared_error = sum(
        _squared_distance(original, decoded)
        for original, decoded in zip(materialized, reconstructed)
    )
    element_count = len(materialized) * quantizer.dimensions
    cosine_sum = sum(
        _cosine(_normalize(original), _normalize(decoded))
        for original, decoded in zip(materialized, reconstructed)
    )
    original_bytes = element_count * float_bytes
    encoded_bytes = len(materialized) * quantizer.encoded_bytes_per_vector
    total_compressed = encoded_bytes + quantizer.codebook_bytes
    ratio = original_bytes / total_compressed if total_compressed else math.inf
    return QuantizationAudit(
        vector_count=len(materialized),
        dimensions=quantizer.dimensions,
        original_bytes=original_bytes,
        encoded_bytes=encoded_bytes,
        codebook_bytes=quantizer.codebook_bytes,
        compression_ratio=ratio,
        mean_squared_error=squared_error / element_count,
        mean_cosine_similarity=cosine_sum / len(materialized),
    )


def recall_at_k(
    exact: Sequence[Union[str, SearchResult]],
    approximate: Sequence[Union[str, SearchResult]],
    k: int,
) -> float:
    """Compute ANN recall as overlap with the exact top-k neighbor ids."""

    if k <= 0:
        raise ValueError("k must be positive")

    def identifier(value: Union[str, SearchResult]) -> str:
        return value if isinstance(value, str) else value.chunk.id

    exact_ids = {identifier(value) for value in exact[:k]}
    approximate_ids = {identifier(value) for value in approximate[:k]}
    return len(exact_ids & approximate_ids) / len(exact_ids) if exact_ids else 1.0


@dataclass(frozen=True)
class ANNRecallAudit:
    """Per-query and mean recall for an approximate index configuration."""

    k: int
    nprobe: int
    per_query: Tuple[float, ...]
    mean_recall: float


def evaluate_ann_recall(
    exact_index: ExactCosineIndex,
    approximate_index: IVFCoarseIndex,
    query_vectors: Sequence[Sequence[float]],
    k: int = 10,
    nprobe: int = 1,
) -> ANNRecallAudit:
    """Compare IVF results with exhaustive cosine neighbors for each query."""

    recalls = tuple(
        recall_at_k(
            exact_index.search(query, k),
            approximate_index.search(query, k, nprobe=nprobe),
            k,
        )
        for query in query_vectors
    )
    mean_recall = sum(recalls) / len(recalls) if recalls else 1.0
    return ANNRecallAudit(k=k, nprobe=nprobe, per_query=recalls, mean_recall=mean_recall)


# Familiar aliases for readers coming from search-engine and ANN terminology.
BM25InvertedIndex = InvertedIndex
IVFIndex = IVFCoarseIndex
ann_recall_at_k = recall_at_k
quantization_audit = audit_quantization
