1"""String distance evaluator."""
2
3from __future__ import annotations
4
5from lexigram.ai.evaluation.evaluators.base import BaseEvaluator
6from lexigram.contracts.ai.evaluation import (
7 EvaluationResult,
8 EvaluationScoreType,
9 EvaluatorProtocol,
10)
11from lexigram.logging import get_logger
12from lexigram.result import Ok, Result
13
14logger = get_logger(__name__)
15
16
17class StringDistanceEvaluator(BaseEvaluator, EvaluatorProtocol):
18 """String similarity evaluation.
19
20 Evaluates output against reference using string distance metrics.
21 Supports Levenshtein, Jaccard, and cosine similarity.
22 """
23
24 def __init__(self, metric: str = "levenshtein") -> None:
25 super().__init__(EvaluationScoreType.STRING_DISTANCE)
26 self._metric = metric
27
28 @property
29 def name(self) -> str:
30 return "string_distance"
31
32 async def evaluate(
33 self,
34 input: str,
35 output: str,
36 reference: str,
37 ) -> Result[EvaluationResult, Exception]:
38 output_norm = output.strip().lower()
39 reference_norm = reference.strip().lower()
40
41 if self._metric == "levenshtein":
42 distance = self._levenshtein_distance(output_norm, reference_norm)
43 max_dist = max(len(output_norm), len(reference_norm))
44 score = 1.0 - (distance / max_dist) if max_dist > 0 else 1.0
45 details: dict[str, float | str] = {
46 "distance": distance,
47 "max_distance": max_dist,
48 }
49 elif self._metric == "jaccard":
50 score = self._jaccard_similarity(output_norm, reference_norm)
51 details = {"metric": "jaccard"}
52 else:
53 score = self._jaccard_similarity(output_norm, reference_norm)
54 details = {"metric": "jaccard"}
55
56 feedback = f"Similarity score: {score:.2f}"
57
58 return Ok(self._create_result(score, feedback, details))
59
60 def _levenshtein_distance(self, s1: str, s2: str) -> int:
61 if len(s1) < len(s2):
62 return self._levenshtein_distance(s2, s1)
63
64 if len(s2) == 0:
65 return len(s1)
66
67 previous_row = list(range(len(s2) + 1))
68 for i, c1 in enumerate(s1):
69 current_row = [i + 1]
70 for j, c2 in enumerate(s2):
71 insertions = previous_row[j + 1] + 1
72 deletions = current_row[j] + 1
73 substitutions = previous_row[j] + (c1 != c2)
74 current_row.append(min(insertions, deletions, substitutions))
75 previous_row = current_row
76
77 return previous_row[-1]
78
79 def _jaccard_similarity(self, s1: str, s2: str) -> float:
80 set1 = set(s1.split())
81 set2 = set(s2.split())
82 if not set1 or not set2:
83 return 0.0
84 intersection = len(set1 & set2)
85 union = len(set1 | set2)
86 return intersection / union if union > 0 else 0.0
87
88
89__all__ = ["StringDistanceEvaluator"]