Skip to main content

llm_token_visualizer/
data.rs

1use serde::{Deserialize, Serialize};
2
3#[derive(Debug, Clone, Serialize, Deserialize)]
4/// One token for the visualizer: its text and a confidence from 0 to 1.
5pub struct TokenInfo {
6    /// Token text, including any leading space.
7    pub text: String,
8    /// Confidence from 0 (unsure) to 1 (sure); for detect mode this is the token probability.
9    pub confidence: f64,
10}
11
12#[derive(Debug, Clone, Serialize, Deserialize)]
13/// A labelled range of tokens, `start..end` (end exclusive).
14pub struct TokenFlag {
15    /// First token index (inclusive).
16    pub start: usize,
17    /// Last token index (exclusive).
18    pub end: usize,
19    /// Label, e.g. `uncertain` or `fact`.
20    pub flag: String,
21    /// Optional explanation shown with the flag.
22    pub description: Option<String>,
23}
24
25#[derive(Debug, Clone, Serialize, Deserialize)]
26/// Input to the visualizer renderers: tokens with confidences and flags.
27pub struct TokenAnalysis {
28    /// The tokens, in order.
29    pub tokens: Vec<TokenInfo>,
30    /// Flagged ranges over `tokens`.
31    pub flags: Vec<TokenFlag>,
32}
33
34#[derive(Debug, Clone)]
35/// Rendering options for the visualizer renderers.
36pub struct VisualizationConfig {
37    /// Show per-token details.
38    pub verbose: bool,
39    /// Print confidence scores next to tokens.
40    pub show_confidence_scores: bool,
41    /// Print the flags section.
42    pub show_flags: bool,
43}
44
45impl Default for VisualizationConfig {
46    fn default() -> Self {
47        Self {
48            verbose: false,
49            show_confidence_scores: true,
50            show_flags: true,
51        }
52    }
53}
54
55#[derive(Debug, Clone, PartialEq)]
56/// Confidence bucket used for colors.
57pub enum ConfidenceLevel {
58    /// Below 0.3.
59    VeryLow,  // 0.0 - 0.3
60    /// 0.3 to 0.5.
61    Low,      // 0.3 - 0.5
62    /// 0.5 to 0.7.
63    Medium,   // 0.5 - 0.7
64    /// 0.7 to 0.9.
65    High,     // 0.7 - 0.9
66    /// 0.9 and above.
67    VeryHigh, // 0.9 - 1.0
68}
69
70impl From<f64> for ConfidenceLevel {
71    fn from(confidence: f64) -> Self {
72        match confidence {
73            c if c < 0.3 => ConfidenceLevel::VeryLow,
74            c if c < 0.5 => ConfidenceLevel::Low,
75            c if c < 0.7 => ConfidenceLevel::Medium,
76            c if c < 0.9 => ConfidenceLevel::High,
77            _ => ConfidenceLevel::VeryHigh,
78        }
79    }
80}
81
82#[derive(Debug, Clone, PartialEq)]
83/// Kind of a [`TokenFlag`], parsed from its label.
84pub enum FlagType {
85    /// `fact`
86    Fact,
87    /// `uncertain`
88    Uncertain,
89    /// `overconfident`
90    Overconfident,
91    /// `hallucination`
92    Hallucination,
93    /// Any other label.
94    Other(String),
95}
96
97impl From<&str> for FlagType {
98    fn from(s: &str) -> Self {
99        match s.to_lowercase().as_str() {
100            "fact" => FlagType::Fact,
101            "uncertain" => FlagType::Uncertain,
102            "overconfident" => FlagType::Overconfident,
103            "hallucination" => FlagType::Hallucination,
104            _ => FlagType::Other(s.to_string()),
105        }
106    }
107}
108
109impl TokenAnalysis {
110    /// Check that there are tokens, every confidence is within 0..=1, and every flag range is inside the token list.
111    pub fn validate(&self) -> Result<(), String> {
112        if self.tokens.is_empty() {
113            return Err("No tokens provided".to_string());
114        }
115
116        for (i, token) in self.tokens.iter().enumerate() {
117            if token.confidence < 0.0 || token.confidence > 1.0 {
118                return Err(format!(
119                    "Invalid confidence score for token {}: {}",
120                    i, token.confidence
121                ));
122            }
123        }
124
125        for flag in &self.flags {
126            if flag.start >= self.tokens.len() || flag.end > self.tokens.len() {
127                return Err(format!(
128                    "Flag span out of bounds: {} to {}",
129                    flag.start, flag.end
130                ));
131            }
132            if flag.start >= flag.end {
133                return Err(format!("Invalid flag span: {} to {}", flag.start, flag.end));
134            }
135        }
136
137        Ok(())
138    }
139
140    /// Flags whose range covers `token_index`.
141    pub fn get_flags_for_token(&self, token_index: usize) -> Vec<&TokenFlag> {
142        self.flags
143            .iter()
144            .filter(|flag| token_index >= flag.start && token_index < flag.end)
145            .collect()
146    }
147
148    /// `(min, max, mean)` confidence; `(0.0, 0.0, 0.0)` when there are no tokens.
149    pub fn get_confidence_stats(&self) -> (f64, f64, f64) {
150        let confidences: Vec<f64> = self.tokens.iter().map(|t| t.confidence).collect();
151        if confidences.is_empty() {
152            // Before 0.5.0 this returned (inf, -inf, NaN).
153            return (0.0, 0.0, 0.0);
154        }
155        let sum: f64 = confidences.iter().sum();
156        let avg = sum / confidences.len() as f64;
157        let min = confidences.iter().fold(f64::INFINITY, |a, &b| a.min(b));
158        let max = confidences.iter().fold(f64::NEG_INFINITY, |a, &b| a.max(b));
159        (min, max, avg)
160    }
161}
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166
167    #[test]
168    fn stats_of_empty_analysis_are_zero_not_nan() {
169        let a = TokenAnalysis { tokens: vec![], flags: vec![] };
170        assert_eq!(a.get_confidence_stats(), (0.0, 0.0, 0.0));
171    }
172}