1use serde::{Deserialize, Serialize};
5
6#[derive(Debug, Clone, Serialize, Deserialize)]
8pub struct EvalMetrics {
9 pub fabricated_entity_rate: f64,
11 pub statement_precision: f64,
13 pub statement_recall: f64,
15 pub statement_f1: f64,
17}
18
19#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
21pub struct ExtractedTriple {
22 pub predicate: String,
23 pub subject_surface: String,
24 pub object_surface: String,
25}
26
27pub fn compute_metrics(extracted: &[ExtractedTriple], expected: &[ExtractedTriple]) -> EvalMetrics {
32 let extracted_set: std::collections::HashSet<(String, String, String)> = extracted
33 .iter()
34 .map(|t| {
35 (
36 t.predicate.to_lowercase(),
37 t.subject_surface.to_lowercase(),
38 t.object_surface.to_lowercase(),
39 )
40 })
41 .collect();
42
43 let expected_set: std::collections::HashSet<(String, String, String)> = expected
44 .iter()
45 .map(|t| {
46 (
47 t.predicate.to_lowercase(),
48 t.subject_surface.to_lowercase(),
49 t.object_surface.to_lowercase(),
50 )
51 })
52 .collect();
53
54 let correct = extracted_set.intersection(&expected_set).count();
55 let extracted_count = extracted_set.len();
56 let expected_count = expected_set.len();
57
58 let precision = if extracted_count > 0 {
59 correct as f64 / extracted_count as f64
60 } else {
61 0.0
62 };
63 let recall = if expected_count > 0 {
64 correct as f64 / expected_count as f64
65 } else {
66 1.0 };
68 let f1 = if precision + recall > 0.0 {
69 2.0 * precision * recall / (precision + recall)
70 } else {
71 0.0
72 };
73
74 let fabricated_entity_rate = 0.0;
77
78 EvalMetrics {
79 fabricated_entity_rate,
80 statement_precision: precision,
81 statement_recall: recall,
82 statement_f1: f1,
83 }
84}
85
86pub fn measure_fabrication_rate(entity_surfaces: &[String], source_text: &str) -> f64 {
92 if entity_surfaces.is_empty() {
93 return 0.0;
94 }
95 let fabricated = entity_surfaces
96 .iter()
97 .filter(|surface| !source_text.contains(surface.as_str()))
98 .count();
99 fabricated as f64 / entity_surfaces.len() as f64
100}
101
102pub fn compute_metrics_with_fabrication(
106 extracted: &[ExtractedTriple],
107 expected: &[ExtractedTriple],
108 fabricated_entity_rate: f64,
109) -> EvalMetrics {
110 let mut metrics = compute_metrics(extracted, expected);
111 metrics.fabricated_entity_rate = fabricated_entity_rate;
112 metrics
113}
114
115impl EvalMetrics {
116 pub fn check_gates(&self) -> Result<(), String> {
118 if self.fabricated_entity_rate != 0.0 {
119 return Err(format!(
120 "fabricated_entity_rate = {:.4}, expected 0.00 (structural hard zero)",
121 self.fabricated_entity_rate
122 ));
123 }
124 if self.statement_precision < 0.90 {
125 return Err(format!(
126 "statement_precision = {:.4}, expected ≥ 0.90",
127 self.statement_precision
128 ));
129 }
130 if self.statement_recall < 0.70 {
131 return Err(format!(
132 "statement_recall = {:.4}, expected ≥ 0.70",
133 self.statement_recall
134 ));
135 }
136 Ok(())
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143
144 #[test]
145 fn perfect_extraction() {
146 let triples = vec![
147 triple("works_on", "Alice", "ProjectX"),
148 triple("employed_by", "Alice", "Acme"),
149 ];
150 let m = compute_metrics(&triples, &triples);
151 assert_eq!(m.statement_precision, 1.0);
152 assert_eq!(m.statement_recall, 1.0);
153 assert_eq!(m.statement_f1, 1.0);
154 assert!(m.check_gates().is_ok());
155 }
156
157 #[test]
158 fn partial_match() {
159 let extracted = vec![
160 triple("works_on", "Alice", "ProjectX"),
161 triple("employed_by", "Alice", "Acme"),
162 ];
163 let expected = vec![
164 triple("works_on", "Alice", "ProjectX"),
165 triple("employed_by", "Alice", "Acme"),
166 triple("knows", "Alice", "Bob"),
167 ];
168 let m = compute_metrics(&extracted, &expected);
169 assert_eq!(m.statement_precision, 1.0);
171 assert!((m.statement_recall - 0.667).abs() < 0.01);
172 }
173
174 #[test]
175 fn case_insensitive_matching() {
176 let extracted = vec![triple("WORKS_ON", "alice", "projectx")];
177 let expected = vec![triple("works_on", "Alice", "ProjectX")];
178 let m = compute_metrics(&extracted, &expected);
179 assert_eq!(m.statement_precision, 1.0);
180 assert_eq!(m.statement_recall, 1.0);
181 }
182
183 #[test]
184 fn empty_extraction_zero_precision() {
185 let expected = vec![triple("works_on", "Alice", "ProjectX")];
186 let m = compute_metrics(&[], &expected);
187 assert_eq!(m.statement_precision, 0.0);
188 assert_eq!(m.statement_recall, 0.0);
189 }
190
191 fn triple(p: &str, s: &str, o: &str) -> ExtractedTriple {
192 ExtractedTriple {
193 predicate: p.into(),
194 subject_surface: s.into(),
195 object_surface: o.into(),
196 }
197 }
198
199 #[test]
202 fn fabrication_rate_zero_when_all_surfaces_present() {
203 let surfaces = vec!["Alice".to_string(), "Acme".to_string()];
204 let text = "Alice works at Acme Corp";
205 assert_eq!(measure_fabrication_rate(&surfaces, text), 0.0);
206 }
207
208 #[test]
209 fn fabrication_rate_nonzero_when_surface_missing() {
210 let surfaces = vec!["Alice".to_string(), "Hallucinated".to_string()];
211 let text = "Alice works at Acme Corp";
212 let rate = measure_fabrication_rate(&surfaces, text);
213 assert!((rate - 0.5).abs() < 0.01, "1 of 2 fabricated, got {rate}");
214 }
215
216 #[test]
217 fn fabrication_rate_empty_input_is_zero() {
218 assert_eq!(measure_fabrication_rate(&[], "any text"), 0.0);
219 }
220}