1use std::fmt::{self, Display};
2
3use serde::{Deserialize, Serialize};
4
5pub type ExtractionRateMetric = crate::eval::governance::RateMetric;
6
7pub const DEFAULT_CORPUS_PATH: &str = "eval/extraction/corpus.json";
8pub const DEFAULT_BASELINE_PATH: &str = "eval/extraction/baseline.json";
9
10#[derive(Debug, Clone)]
11pub struct ExtractionEvalOptions {
12 pub corpus_path: String,
13}
14
15impl Default for ExtractionEvalOptions {
16 fn default() -> Self {
17 Self {
18 corpus_path: DEFAULT_CORPUS_PATH.to_string(),
19 }
20 }
21}
22
23#[derive(Debug, Clone, Deserialize)]
24pub(crate) struct ExtractionCorpus {
25 pub version: String,
26 #[serde(default)]
27 pub description: String,
28 pub cases: Vec<ExtractionCase>,
29}
30
31#[derive(Debug, Clone, Deserialize)]
32pub(crate) struct ExtractionCase {
33 pub id: String,
34 #[serde(default)]
35 pub transcript: Vec<TranscriptEvent>,
36 pub observation_output: String,
37 pub candidate_output: String,
38 #[serde(default)]
39 pub expected_observations: Vec<ObservationExpectation>,
40 #[serde(default)]
41 pub forbidden_observations: Vec<ObservationExpectation>,
42 #[serde(default)]
43 pub expected_candidates: Vec<CandidateExpectation>,
44 #[serde(default)]
45 pub forbidden_candidates: Vec<CandidateExpectation>,
46}
47
48#[derive(Debug, Clone, Deserialize)]
49pub(crate) struct TranscriptEvent {
50 pub id: String,
51 pub role: String,
52 pub content: String,
53 #[serde(default)]
54 pub tool_name: Option<String>,
55 #[serde(default)]
56 pub event_type: Option<String>,
57 #[serde(default)]
58 pub token_estimate: Option<i64>,
59 #[serde(default)]
60 pub created_at_epoch: Option<i64>,
61}
62
63#[derive(Debug, Clone, Deserialize)]
64pub(crate) struct ObservationExpectation {
65 pub id: String,
66 #[serde(default)]
67 pub observation_type: Option<String>,
68 #[serde(default)]
69 pub text_contains: Vec<String>,
70}
71
72#[derive(Debug, Clone, Deserialize)]
73pub(crate) struct CandidateExpectation {
74 pub id: String,
75 #[serde(default)]
76 pub scope: Option<String>,
77 #[serde(default)]
78 pub memory_type: Option<String>,
79 #[serde(default)]
80 pub topic_key: Option<String>,
81 #[serde(default)]
82 pub risk_class: Option<String>,
83 #[serde(default)]
84 pub text_contains: Vec<String>,
85}
86
87#[derive(Debug, Clone, Serialize, PartialEq)]
88pub struct ExtractionEvalReport {
89 pub metadata: ExtractionEvalMetadata,
90 pub metrics: ExtractionMetricSummary,
91 pub cases: Vec<ExtractionCaseReport>,
92 pub failing_examples: Vec<String>,
93}
94
95#[derive(Debug, Clone, Serialize, PartialEq)]
96pub struct ExtractionEvalMetadata {
97 pub corpus: String,
98 pub corpus_version: String,
99 pub description: String,
100 pub cases: usize,
101 pub transcript_events: usize,
102}
103
104#[derive(Debug, Clone, Serialize, PartialEq)]
105pub struct ExtractionMetricSummary {
106 pub observation_precision: ExtractionRateMetric,
107 pub observation_recall: ExtractionRateMetric,
108 pub candidate_precision: ExtractionRateMetric,
109 pub candidate_recall: ExtractionRateMetric,
110 pub forbidden_observation_exclusion: ExtractionRateMetric,
111 pub forbidden_candidate_exclusion: ExtractionRateMetric,
112 pub candidate_risk_classes: CandidateRiskClassCounts,
113 pub over_saved_predictions: usize,
114 pub total_predictions: usize,
115 pub over_save_penalty: f64,
116 pub all_checks_passed: bool,
117}
118
119#[derive(Debug, Clone, Default, Serialize, PartialEq, Eq)]
120pub struct CandidateRiskClassCounts {
121 pub low: usize,
122 pub medium: usize,
123 pub high: usize,
124}
125
126#[derive(Debug, Clone, Serialize, PartialEq)]
127pub struct ExtractionCaseReport {
128 pub id: String,
129 pub transcript_events: usize,
130 pub observation_request_sha256: String,
131 pub candidate_request_sha256: String,
132 pub predicted_observations: Vec<ObservationPrediction>,
133 pub predicted_candidates: Vec<CandidatePrediction>,
134 pub missing_expected_observations: Vec<String>,
135 pub unexpected_observations: Vec<usize>,
136 pub forbidden_observations: Vec<String>,
137 pub missing_expected_candidates: Vec<String>,
138 pub unexpected_candidates: Vec<usize>,
139 pub forbidden_candidates: Vec<String>,
140 pub over_saved_predictions: usize,
141 pub pass: bool,
142}
143
144#[derive(Debug, Clone, Serialize, PartialEq)]
145pub struct ObservationPrediction {
146 pub index: usize,
147 pub observation_type: String,
148 pub text: String,
149 pub confidence: Option<f64>,
150}
151
152#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
153pub struct CandidatePrediction {
154 pub index: usize,
155 pub scope: String,
156 pub memory_type: String,
157 pub topic_key: String,
158 pub risk_class: String,
159 pub text: String,
160}
161
162impl Display for ExtractionEvalReport {
163 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
164 writeln!(
165 f,
166 "Extraction eval: {} cases from {}",
167 self.metadata.cases, self.metadata.corpus
168 )?;
169 writeln!(
170 f,
171 "Observations: precision {} recall {} forbidden exclusion {}",
172 format_rate(&self.metrics.observation_precision),
173 format_rate(&self.metrics.observation_recall),
174 format_rate(&self.metrics.forbidden_observation_exclusion)
175 )?;
176 writeln!(
177 f,
178 "Candidates: precision {} recall {} forbidden exclusion {}",
179 format_rate(&self.metrics.candidate_precision),
180 format_rate(&self.metrics.candidate_recall),
181 format_rate(&self.metrics.forbidden_candidate_exclusion)
182 )?;
183 writeln!(
184 f,
185 "Candidate risk classes: low={} medium={} high={}",
186 self.metrics.candidate_risk_classes.low,
187 self.metrics.candidate_risk_classes.medium,
188 self.metrics.candidate_risk_classes.high
189 )?;
190 writeln!(
191 f,
192 "Over-save penalty: {:.4} ({}/{})",
193 self.metrics.over_save_penalty,
194 self.metrics.over_saved_predictions,
195 self.metrics.total_predictions
196 )?;
197 writeln!(f, "All checks passed: {}", self.metrics.all_checks_passed)?;
198 if !self.failing_examples.is_empty() {
199 writeln!(f, "Failures:")?;
200 for failure in &self.failing_examples {
201 writeln!(f, "- {failure}")?;
202 }
203 }
204 Ok(())
205 }
206}
207
208fn format_rate(metric: &ExtractionRateMetric) -> String {
209 format!("{}/{} ({:.4})", metric.passed, metric.total, metric.rate)
210}