"""Pure retrieval metrics with explicit label-contract validation.""" from math import log2 from typing import Sequence, Union def _validate(relevances: Sequence[Union[int, float]], k: int) -> None: if k <= 0: raise ValueError("k must be positive") if any(value < 0 for value in relevances): raise ValueError("relevance values must be non-negative") def recall_at_k( relevances: Sequence[Union[int, float]], k: int, total_relevant: int, ) -> float: _validate(relevances, k) if total_relevant <= 0: raise ValueError("total_relevant must be positive") retrieved = sum(value > 0 for value in relevances[:k]) if retrieved > total_relevant: raise ValueError("retrieved labels exceed total_relevant") return retrieved / total_relevant def reciprocal_rank(relevances: Sequence[Union[int, float]]) -> float: for rank, value in enumerate(relevances, start=1): if value > 0: return 1 / rank return 0.0 def dcg_at_k(relevances: Sequence[Union[int, float]], k: int) -> float: _validate(relevances, k) return sum( (2**relevance - 1) / log2(rank + 1) for rank, relevance in enumerate(relevances[:k], start=1) ) def ndcg_at_k(relevances: Sequence[Union[int, float]], k: int) -> float: _validate(relevances, k) ideal = dcg_at_k(sorted(relevances, reverse=True), k) return 0.0 if ideal == 0 else dcg_at_k(relevances, k) / ideal