use crate::TokenEvent;
use std::io::Write;
pub struct HeatmapExporter {
data: Vec<Vec<Option<f32>>>,
run_count: usize,
}
impl HeatmapExporter {
pub fn new() -> Self {
HeatmapExporter {
data: Vec::new(),
run_count: 0,
}
}
pub fn record_run(&mut self, events: &[TokenEvent]) {
let run_idx = self.run_count;
self.run_count += 1;
if events.len() > self.data.len() {
self.data.resize(events.len(), Vec::new());
}
for (pos, event) in events.iter().enumerate() {
while self.data[pos].len() < run_idx {
self.data[pos].push(None);
}
self.data[pos].push(event.confidence);
}
}
pub fn export_csv(
&self,
path: &str,
min_confidence: f32,
sort_by: &str,
) -> Result<(), Box<dyn std::error::Error>> {
let mut file = std::fs::File::create(path)?;
let header_cols: Vec<String> = (0..self.run_count).map(|i| format!("run_{}", i)).collect();
writeln!(file, "position,{}", header_cols.join(","))?;
let mut rows: Vec<(usize, f32, String)> = self
.data
.iter()
.enumerate()
.map(|(pos, runs)| {
let vals: Vec<f32> = runs.iter().filter_map(|v| *v).collect();
let mean = if vals.is_empty() {
0.0f32
} else {
vals.iter().sum::<f32>() / vals.len() as f32
};
let cols: Vec<String> = (0..self.run_count)
.map(|r| {
runs.get(r)
.and_then(|v| *v)
.map(|v| v.to_string())
.unwrap_or_default()
})
.collect();
(pos, mean, format!("{},{}", pos, cols.join(",")))
})
.filter(|(_, mean, _)| *mean >= min_confidence)
.collect();
if sort_by == "confidence" {
rows.sort_by(|a, b| {
match (a.1.is_nan(), b.1.is_nan()) {
(true, true) => std::cmp::Ordering::Equal,
(true, false) => std::cmp::Ordering::Greater, (false, true) => std::cmp::Ordering::Less,
(false, false) => b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal),
}
});
}
for (_, _, row) in &rows {
writeln!(file, "{}", row)?;
}
Ok(())
}
}
impl Default for HeatmapExporter {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::TokenEvent;
fn make_event(idx: usize, confidence: Option<f32>) -> TokenEvent {
TokenEvent {
text: "tok".to_string(),
original: "tok".to_string(),
index: idx,
transformed: false,
importance: 0.0,
chaos_label: None,
provider: None,
confidence,
perplexity: None,
alternatives: vec![],
is_error: false,
arrival_ms: None,
}
}
#[test]
fn test_record_and_export() {
let mut exporter = HeatmapExporter::new();
let events = vec![make_event(0, Some(0.9)), make_event(1, Some(0.8))];
exporter.record_run(&events);
let tmp = std::env::temp_dir().join("heatmap_test.csv");
exporter
.export_csv(tmp.to_str().unwrap(), 0.0, "position")
.expect("export");
let content = std::fs::read_to_string(&tmp).expect("read");
assert!(content.contains("position,run_0"));
assert!(content.contains("0,0.9"));
std::fs::remove_file(&tmp).ok();
}
#[test]
fn test_empty_exporter() {
let exporter = HeatmapExporter::new();
let tmp = std::env::temp_dir().join("heatmap_empty.csv");
exporter
.export_csv(tmp.to_str().unwrap(), 0.0, "position")
.expect("export empty");
std::fs::remove_file(&tmp).ok();
}
#[test]
fn test_nan_confidence_sorts_last() {
let mut exporter = HeatmapExporter::new();
let events = vec![
make_event(0, Some(0.9)),
make_event(1, None), make_event(2, Some(0.5)),
];
exporter.record_run(&events);
let tmp = std::env::temp_dir().join("heatmap_nan_test.csv");
exporter
.export_csv(tmp.to_str().unwrap(), 0.0, "confidence")
.expect("export");
let content = std::fs::read_to_string(&tmp).expect("read");
std::fs::remove_file(&tmp).ok();
let lines: Vec<&str> = content.lines().collect();
assert!(lines.len() >= 3, "should have header + 3 data rows");
let last_row = lines.last().unwrap();
assert!(last_row.starts_with("1,"), "NaN-confidence row should sort last");
}
#[test]
fn test_min_confidence_filter() {
let mut exporter = HeatmapExporter::new();
let events = vec![
make_event(0, Some(0.9)),
make_event(1, Some(0.1)),
];
exporter.record_run(&events);
let tmp = std::env::temp_dir().join("heatmap_filter_test.csv");
exporter
.export_csv(tmp.to_str().unwrap(), 0.5, "position")
.expect("export");
let content = std::fs::read_to_string(&tmp).expect("read");
std::fs::remove_file(&tmp).ok();
assert!(content.contains("0,0.9"), "high-confidence row should be present");
assert!(!content.contains("1,0.1"), "low-confidence row should be filtered");
}
#[test]
fn test_multiple_runs_alignment() {
let mut exporter = HeatmapExporter::new();
exporter.record_run(&[make_event(0, Some(0.8)), make_event(1, Some(0.7))]);
exporter.record_run(&[make_event(0, Some(0.6))]);
let tmp = std::env::temp_dir().join("heatmap_multirun.csv");
exporter
.export_csv(tmp.to_str().unwrap(), 0.0, "position")
.expect("export");
let content = std::fs::read_to_string(&tmp).expect("read");
std::fs::remove_file(&tmp).ok();
assert!(content.contains("run_0"));
assert!(content.contains("run_1"));
}
}