Skip to main content

remem/eval/extraction/
types.rs

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}