Skip to main content

llm_token_visualizer/
detect.rs

1//! Hallucination-risk detection from token log probabilities.
2//!
3//! The input is the `logprobs` block that OpenAI-compatible Chat Completions
4//! APIs return when you ask for `"logprobs": true`. Each token's probability
5//! is `exp(logprob)`. Words containing a token whose probability is below the
6//! threshold are flagged, and neighbouring flagged words are merged into one
7//! span. Each span reports the weakest token and the alternatives the model
8//! was weighing at that point.
9//!
10//! Low probability is a signal, not proof: a model can be confidently wrong,
11//! and it can be unsure about phrasing while the facts are right.
12
13use anyhow::{bail, Context, Result};
14use serde::{Deserialize, Serialize};
15use serde_json::Value;
16
17use crate::data::{TokenAnalysis, TokenFlag, TokenInfo};
18
19/// Default probability below which a token counts as low confidence.
20pub const DEFAULT_THRESHOLD: f64 = 0.5;
21
22/// Label used for spans produced by [`detect`].
23pub const FLAG_LABEL: &str = "uncertain";
24
25#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
26pub struct Alternative {
27    pub token: String,
28    pub logprob: f64,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
32pub struct LogprobToken {
33    pub token: String,
34    pub logprob: f64,
35    #[serde(default)]
36    pub top_logprobs: Vec<Alternative>,
37}
38
39impl LogprobToken {
40    pub fn prob(&self) -> f64 {
41        self.logprob.exp().clamp(0.0, 1.0)
42    }
43}
44
45/// A run of low-confidence words.
46#[derive(Debug, Clone, Serialize, PartialEq)]
47pub struct Span {
48    /// First token index (inclusive).
49    pub start: usize,
50    /// Last token index (exclusive).
51    pub end: usize,
52    pub text: String,
53    /// The least likely token in the span and its probability.
54    pub weakest_token: String,
55    pub min_prob: f64,
56    /// What the model considered at the weakest token, most likely first.
57    /// Empty when the input has no `top_logprobs`.
58    pub alternatives: Vec<Candidate>,
59}
60
61#[derive(Debug, Clone, Serialize, PartialEq)]
62pub struct Candidate {
63    pub token: String,
64    pub prob: f64,
65}
66
67#[derive(Debug, Clone, Serialize, PartialEq)]
68pub struct Report {
69    pub threshold: f64,
70    pub text: String,
71    pub token_count: usize,
72    pub mean_prob: f64,
73    pub spans: Vec<Span>,
74}
75
76impl Report {
77    pub fn flagged(&self) -> bool {
78        !self.spans.is_empty()
79    }
80}
81
82/// Extract the token list from any of these shapes:
83/// a full Chat Completions response (`choices[0].logprobs.content`),
84/// a `logprobs` object (`{"content": [...]}`), or a bare token array.
85pub fn parse_logprobs(json: &str) -> Result<Vec<LogprobToken>> {
86    let value: Value = serde_json::from_str(json).context("input is not valid JSON")?;
87    if let Some(err) = value.get("error") {
88        bail!("the API returned an error: {}", err);
89    }
90    let content = if value.is_array() {
91        &value
92    } else if let Some(c) = value.pointer("/choices/0/logprobs/content") {
93        c
94    } else if let Some(c) = value.get("content") {
95        c
96    } else {
97        bail!(
98            "no token logprobs found; expected choices[0].logprobs.content \
99             (request the completion with \"logprobs\": true)"
100        );
101    };
102    if content.is_null() {
103        bail!("logprobs are null; request the completion with \"logprobs\": true");
104    }
105    let mut tokens: Vec<LogprobToken> =
106        serde_json::from_value(content.clone()).context("could not read logprobs tokens")?;
107    // Drop end-of-turn and other special tokens (e.g. "<|eot_id|>") that some
108    // servers include in the logprobs but not in the message text.
109    tokens.retain(|t| !is_special_token(&t.token));
110    if tokens.is_empty() {
111        bail!("logprobs contain no tokens");
112    }
113    Ok(tokens)
114}
115
116fn is_special_token(s: &str) -> bool {
117    s.starts_with("<|") && s.ends_with("|>")
118}
119
120fn has_word_chars(s: &str) -> bool {
121    s.chars().any(|c| c.is_alphanumeric())
122}
123
124/// A token starts a new word if it begins with whitespace or punctuation,
125/// or if the previous token ended in one (e.g. " " followed by "198").
126fn starts_word(tokens: &[LogprobToken], i: usize) -> bool {
127    if i == 0 {
128        return true;
129    }
130    let first = tokens[i].token.chars().next();
131    let prev_last = tokens[i - 1].token.chars().last();
132    let boundary =
133        |c: Option<char>| c.is_none_or(|c| c.is_whitespace() || c.is_ascii_punctuation());
134    boundary(first) || boundary(prev_last)
135}
136
137/// Group token indices into words: returns (start, end) ranges.
138fn words(tokens: &[LogprobToken]) -> Vec<(usize, usize)> {
139    let mut out: Vec<(usize, usize)> = Vec::new();
140    for i in 0..tokens.len() {
141        if starts_word(tokens, i) || out.is_empty() {
142            out.push((i, i + 1));
143        } else if let Some(last) = out.last_mut() {
144            last.1 = i + 1;
145        }
146    }
147    out
148}
149
150/// Flag low-confidence spans.
151pub fn detect(tokens: &[LogprobToken], threshold: f64) -> Report {
152    let text: String = tokens.iter().map(|t| t.token.as_str()).collect();
153    let mean_prob = if tokens.is_empty() {
154        0.0
155    } else {
156        tokens.iter().map(|t| t.prob()).sum::<f64>() / tokens.len() as f64
157    };
158
159    // A word is flagged when any of its tokens with letters or digits is
160    // below the threshold. Pure whitespace/punctuation never triggers a flag.
161    let flagged_words: Vec<(usize, usize)> = words(tokens)
162        .into_iter()
163        .filter(|&(s, e)| {
164            tokens[s..e]
165                .iter()
166                .any(|t| has_word_chars(&t.token) && t.prob() < threshold)
167        })
168        .collect();
169
170    // Merge words that are adjacent, or separated only by whitespace tokens.
171    let mut ranges: Vec<(usize, usize)> = Vec::new();
172    for (s, e) in flagged_words {
173        if let Some(last) = ranges.last_mut() {
174            if tokens[last.1..s].iter().all(|t| t.token.trim().is_empty()) {
175                last.1 = e;
176                continue;
177            }
178        }
179        ranges.push((s, e));
180    }
181
182    let spans = ranges
183        .into_iter()
184        .map(|(start, end)| {
185            let weakest = (start..end)
186                .filter(|&i| has_word_chars(&tokens[i].token))
187                .min_by(|&a, &b| tokens[a].prob().total_cmp(&tokens[b].prob()))
188                .unwrap_or(start);
189            let w = &tokens[weakest];
190            let alternatives = w
191                .top_logprobs
192                .iter()
193                .map(|a| Candidate {
194                    token: a.token.clone(),
195                    prob: a.logprob.exp(),
196                })
197                .collect();
198            Span {
199                start,
200                end,
201                text: tokens[start..end]
202                    .iter()
203                    .map(|t| t.token.as_str())
204                    .collect(),
205                weakest_token: w.token.clone(),
206                min_prob: w.prob(),
207                alternatives,
208            }
209        })
210        .collect();
211
212    Report {
213        threshold,
214        text,
215        token_count: tokens.len(),
216        mean_prob,
217        spans,
218    }
219}
220
221impl Span {
222    /// One-line human description, e.g.
223    /// `p=0.57 at "ord"; model also considered "üsseldorf" (0.39)`.
224    pub fn describe(&self) -> String {
225        let mut s = format!("p={:.2} at {:?}", self.min_prob, self.weakest_token);
226        let others: Vec<String> = self
227            .alternatives
228            .iter()
229            .filter(|c| c.token != self.weakest_token)
230            .map(|c| format!("{:?} ({:.2})", c.token, c.prob))
231            .collect();
232        if !others.is_empty() {
233            s.push_str("; model also considered ");
234            s.push_str(&others.join(", "));
235        }
236        s
237    }
238}
239
240/// Convert tokens plus a report into the visualizer's input format, so the
241/// terminal, HTML and Markdown renderers can display the result.
242pub fn to_token_analysis(tokens: &[LogprobToken], report: &Report) -> TokenAnalysis {
243    TokenAnalysis {
244        tokens: tokens
245            .iter()
246            .map(|t| TokenInfo {
247                text: t.token.clone(),
248                confidence: t.prob(),
249            })
250            .collect(),
251        flags: report
252            .spans
253            .iter()
254            .map(|s| TokenFlag {
255                start: s.start,
256                end: s.end,
257                flag: FLAG_LABEL.to_string(),
258                description: Some(s.describe()),
259            })
260            .collect(),
261    }
262}
263
264#[cfg(test)]
265mod tests {
266    use super::*;
267
268    fn tok(t: &str, p: f64) -> LogprobToken {
269        LogprobToken {
270            token: t.to_string(),
271            logprob: p.ln(),
272            top_logprobs: vec![],
273        }
274    }
275
276    #[test]
277    fn parses_all_three_shapes() {
278        let arr = r#"[{"token":"Hi","logprob":-0.1}]"#;
279        let obj = r#"{"content":[{"token":"Hi","logprob":-0.1}]}"#;
280        let full = r#"{"choices":[{"logprobs":{"content":[{"token":"Hi","logprob":-0.1,"top_logprobs":[]}]}}]}"#;
281        for s in [arr, obj, full] {
282            let t = parse_logprobs(s).unwrap();
283            assert_eq!(t.len(), 1);
284            assert_eq!(t[0].token, "Hi");
285        }
286    }
287
288    #[test]
289    fn special_tokens_are_dropped() {
290        let t = parse_logprobs(
291            r#"[{"token":"Hi","logprob":-0.1},{"token":"<|eot_id|>","logprob":-3.0}]"#,
292        )
293        .unwrap();
294        assert_eq!(t.len(), 1);
295    }
296
297    #[test]
298    fn missing_or_null_logprobs_is_a_clear_error() {
299        let e = parse_logprobs(r#"{"choices":[{"logprobs":null}]}"#).unwrap_err();
300        assert!(e.to_string().contains("logprobs"));
301        let e = parse_logprobs(r#"{"error":{"message":"bad key"}}"#).unwrap_err();
302        assert!(e.to_string().contains("bad key"));
303    }
304
305    #[test]
306    fn flags_whole_word_when_a_subword_token_is_low() {
307        let t = vec![
308            tok("In", 0.9),
309            tok(" D", 0.95),
310            tok("ord", 0.3),
311            tok("recht", 1.0),
312        ];
313        let r = detect(&t, 0.5);
314        assert_eq!(r.spans.len(), 1);
315        assert_eq!(r.spans[0].text, " Dordrecht");
316        assert_eq!((r.spans[0].start, r.spans[0].end), (1, 4));
317        assert_eq!(r.spans[0].weakest_token, "ord");
318    }
319
320    #[test]
321    fn punctuation_and_whitespace_never_trigger() {
322        let t = vec![
323            tok("Yes", 0.9),
324            tok(",", 0.1),
325            tok(" ", 0.1),
326            tok("ok", 0.9),
327        ];
328        assert!(!detect(&t, 0.5).flagged());
329    }
330
331    #[test]
332    fn adjacent_low_words_merge_into_one_span() {
333        let t = vec![
334            tok("It", 0.99),
335            tok(" was", 0.2),
336            tok(" Pete", 0.3),
337            tok(".", 0.99),
338            tok(" Then", 0.1),
339        ];
340        let r = detect(&t, 0.5);
341        assert_eq!(r.spans.len(), 2);
342        assert_eq!(r.spans[0].text, " was Pete");
343        assert_eq!(r.spans[1].text, " Then");
344    }
345
346    #[test]
347    fn threshold_is_strict_less_than() {
348        let t = vec![tok("a", 0.5)];
349        assert!(!detect(&t, 0.5).flagged());
350        assert!(detect(&t, 0.51).flagged());
351    }
352
353    #[test]
354    fn digits_split_by_space_token_start_a_new_word() {
355        let t = vec![
356            tok(" in", 0.99),
357            tok(" ", 0.99),
358            tok("169", 0.2),
359            tok("1", 0.99),
360        ];
361        let r = detect(&t, 0.5);
362        assert_eq!(r.spans.len(), 1);
363        assert_eq!(r.spans[0].text, "1691");
364    }
365
366    #[test]
367    fn describe_lists_alternatives_but_not_the_chosen_token() {
368        let mut w = tok("ord", 0.57);
369        w.top_logprobs = vec![
370            Alternative {
371                token: "ord".into(),
372                logprob: 0.57f64.ln(),
373            },
374            Alternative {
375                token: "üsseldorf".into(),
376                logprob: 0.39f64.ln(),
377            },
378        ];
379        let r = detect(&[tok(" D", 0.99), w], 0.6);
380        let d = r.spans[0].describe();
381        assert!(d.contains("p=0.57"), "{d}");
382        assert!(d.contains("üsseldorf"), "{d}");
383        assert_eq!(d.matches("\"ord\"").count(), 1, "{d}");
384    }
385
386    #[test]
387    fn bundled_sample_flags_the_birthplace() {
388        let json = include_str!("../examples/logprobs/cuyp.json");
389        let tokens = parse_logprobs(json).unwrap();
390        let r = detect(&tokens, 0.6);
391        assert_eq!(
392            r.text,
393            "Aelbert Cuyp died in 1691 in Dordrecht, Netherlands."
394        );
395        assert!(r.spans.iter().any(|s| s.text == " Dordrecht"));
396        let a = to_token_analysis(&tokens, &r);
397        assert!(a.validate().is_ok());
398    }
399}