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    // Completions-style shape (`tokens` + `token_logprobs`), which some
91    // servers return from Chat Completions too.
92    let legacy = value
93        .pointer("/choices/0/logprobs")
94        .or_else(|| Some(&value).filter(|v| v.get("token_logprobs").is_some()))
95        .filter(|l| l.get("tokens").is_some() && l.get("token_logprobs").is_some());
96    if let Some(l) = legacy {
97        let mut tokens = parse_legacy(l)?;
98        tokens.retain(|t| !is_special_token(&t.token));
99        if tokens.is_empty() {
100            bail!("logprobs contain no tokens");
101        }
102        return Ok(tokens);
103    }
104    let content = if value.is_array() {
105        &value
106    } else if let Some(c) = value.pointer("/choices/0/logprobs/content") {
107        c
108    } else if let Some(c) = value.get("content") {
109        c
110    } else {
111        bail!(
112            "no token logprobs found; expected choices[0].logprobs.content \
113             (request the completion with \"logprobs\": true)"
114        );
115    };
116    if content.is_null() {
117        bail!("logprobs are null; request the completion with \"logprobs\": true");
118    }
119    let mut tokens: Vec<LogprobToken> =
120        serde_json::from_value(content.clone()).context("could not read logprobs tokens")?;
121    // Drop end-of-turn and other special tokens (e.g. "<|eot_id|>") that some
122    // servers include in the logprobs but not in the message text.
123    tokens.retain(|t| !is_special_token(&t.token));
124    if tokens.is_empty() {
125        bail!("logprobs contain no tokens");
126    }
127    Ok(tokens)
128}
129
130/// Read `{"tokens": [..], "token_logprobs": [..], "top_logprobs": [..]}`,
131/// where each `top_logprobs` entry is either a `{token: logprob}` map or a
132/// list of `{"token", "logprob"}` objects.
133fn parse_legacy(l: &Value) -> Result<Vec<LogprobToken>> {
134    let toks = l["tokens"]
135        .as_array()
136        .context("logprobs.tokens is not a list")?;
137    let lps = l["token_logprobs"]
138        .as_array()
139        .context("logprobs.token_logprobs is not a list")?;
140    if toks.len() != lps.len() {
141        bail!("logprobs.tokens and logprobs.token_logprobs have different lengths");
142    }
143    let tops = l.get("top_logprobs").and_then(Value::as_array);
144    let mut out = Vec::with_capacity(toks.len());
145    for (i, (t, lp)) in toks.iter().zip(lps).enumerate() {
146        let token = t
147            .as_str()
148            .context("a logprobs token is not a string")?
149            .to_string();
150        let logprob = lp.as_f64().unwrap_or(f64::NEG_INFINITY);
151        let mut top_logprobs: Vec<Alternative> = match tops.and_then(|a| a.get(i)) {
152            Some(Value::Object(m)) => m
153                .iter()
154                .filter_map(|(k, v)| {
155                    v.as_f64().map(|lp| Alternative {
156                        token: k.clone(),
157                        logprob: lp,
158                    })
159                })
160                .collect(),
161            Some(v @ Value::Array(_)) => serde_json::from_value(v.clone()).unwrap_or_default(),
162            _ => Vec::new(),
163        };
164        top_logprobs.sort_by(|a, b| b.logprob.total_cmp(&a.logprob));
165        out.push(LogprobToken {
166            token,
167            logprob,
168            top_logprobs,
169        });
170    }
171    Ok(out)
172}
173
174fn is_special_token(s: &str) -> bool {
175    s.starts_with("<|") && s.ends_with("|>")
176}
177
178fn has_word_chars(s: &str) -> bool {
179    s.chars().any(|c| c.is_alphanumeric())
180}
181
182/// A token starts a new word if it begins with whitespace or punctuation,
183/// or if the previous token ended in one (e.g. " " followed by "198").
184fn starts_word(tokens: &[LogprobToken], i: usize) -> bool {
185    if i == 0 {
186        return true;
187    }
188    let first = tokens[i].token.chars().next();
189    let prev_last = tokens[i - 1].token.chars().last();
190    let boundary =
191        |c: Option<char>| c.is_none_or(|c| c.is_whitespace() || c.is_ascii_punctuation());
192    boundary(first) || boundary(prev_last)
193}
194
195/// Group token indices into words: returns (start, end) ranges.
196fn words(tokens: &[LogprobToken]) -> Vec<(usize, usize)> {
197    let mut out: Vec<(usize, usize)> = Vec::new();
198    for i in 0..tokens.len() {
199        if starts_word(tokens, i) || out.is_empty() {
200            out.push((i, i + 1));
201        } else if let Some(last) = out.last_mut() {
202            last.1 = i + 1;
203        }
204    }
205    out
206}
207
208/// Flag low-confidence spans.
209pub fn detect(tokens: &[LogprobToken], threshold: f64) -> Report {
210    let text: String = tokens.iter().map(|t| t.token.as_str()).collect();
211    let mean_prob = if tokens.is_empty() {
212        0.0
213    } else {
214        tokens.iter().map(|t| t.prob()).sum::<f64>() / tokens.len() as f64
215    };
216
217    // A word is flagged when any of its tokens with letters or digits is
218    // below the threshold. Pure whitespace/punctuation never triggers a flag.
219    let flagged_words: Vec<(usize, usize)> = words(tokens)
220        .into_iter()
221        .filter(|&(s, e)| {
222            tokens[s..e]
223                .iter()
224                .any(|t| has_word_chars(&t.token) && t.prob() < threshold)
225        })
226        .collect();
227
228    // Merge words that are adjacent, or separated only by whitespace tokens.
229    let mut ranges: Vec<(usize, usize)> = Vec::new();
230    for (s, e) in flagged_words {
231        if let Some(last) = ranges.last_mut() {
232            if tokens[last.1..s].iter().all(|t| t.token.trim().is_empty()) {
233                last.1 = e;
234                continue;
235            }
236        }
237        ranges.push((s, e));
238    }
239
240    let spans = ranges
241        .into_iter()
242        .map(|(start, end)| {
243            let weakest = (start..end)
244                .filter(|&i| has_word_chars(&tokens[i].token))
245                .min_by(|&a, &b| tokens[a].prob().total_cmp(&tokens[b].prob()))
246                .unwrap_or(start);
247            let w = &tokens[weakest];
248            let alternatives = w
249                .top_logprobs
250                .iter()
251                .map(|a| Candidate {
252                    token: a.token.clone(),
253                    prob: a.logprob.exp(),
254                })
255                .collect();
256            Span {
257                start,
258                end,
259                text: tokens[start..end]
260                    .iter()
261                    .map(|t| t.token.as_str())
262                    .collect(),
263                weakest_token: w.token.clone(),
264                min_prob: w.prob(),
265                alternatives,
266            }
267        })
268        .collect();
269
270    Report {
271        threshold,
272        text,
273        token_count: tokens.len(),
274        mean_prob,
275        spans,
276    }
277}
278
279impl Span {
280    /// One-line human description, e.g.
281    /// `p=0.57 at "ord"; model also considered "üsseldorf" (0.39)`.
282    pub fn describe(&self) -> String {
283        let mut s = format!("p={:.2} at {:?}", self.min_prob, self.weakest_token);
284        let others: Vec<String> = self
285            .alternatives
286            .iter()
287            .filter(|c| c.token != self.weakest_token)
288            .map(|c| format!("{:?} ({:.2})", c.token, c.prob))
289            .collect();
290        if !others.is_empty() {
291            s.push_str("; model also considered ");
292            s.push_str(&others.join(", "));
293        }
294        s
295    }
296}
297
298/// Convert tokens plus a report into the visualizer's input format, so the
299/// terminal, HTML and Markdown renderers can display the result.
300pub fn to_token_analysis(tokens: &[LogprobToken], report: &Report) -> TokenAnalysis {
301    TokenAnalysis {
302        tokens: tokens
303            .iter()
304            .map(|t| TokenInfo {
305                text: t.token.clone(),
306                confidence: t.prob(),
307            })
308            .collect(),
309        flags: report
310            .spans
311            .iter()
312            .map(|s| TokenFlag {
313                start: s.start,
314                end: s.end,
315                flag: FLAG_LABEL.to_string(),
316                description: Some(s.describe()),
317            })
318            .collect(),
319    }
320}
321
322#[cfg(test)]
323mod tests {
324    use super::*;
325
326    fn tok(t: &str, p: f64) -> LogprobToken {
327        LogprobToken {
328            token: t.to_string(),
329            logprob: p.ln(),
330            top_logprobs: vec![],
331        }
332    }
333
334    #[test]
335    fn parses_all_three_shapes() {
336        let arr = r#"[{"token":"Hi","logprob":-0.1}]"#;
337        let obj = r#"{"content":[{"token":"Hi","logprob":-0.1}]}"#;
338        let full = r#"{"choices":[{"logprobs":{"content":[{"token":"Hi","logprob":-0.1,"top_logprobs":[]}]}}]}"#;
339        for s in [arr, obj, full] {
340            let t = parse_logprobs(s).unwrap();
341            assert_eq!(t.len(), 1);
342            assert_eq!(t[0].token, "Hi");
343        }
344    }
345
346    #[test]
347    fn parses_tokens_and_token_logprobs_shape() {
348        let map = r#"{"choices":[{"logprobs":{"tokens":["Hi","!"],"token_logprobs":[-0.1,-2.0],
349            "top_logprobs":[{"Hello":-2.5,"Hi":-0.1},{"!":-2.0}]}}]}"#;
350        let t = parse_logprobs(map).unwrap();
351        assert_eq!(t.len(), 2);
352        assert_eq!(t[0].token, "Hi");
353        assert_eq!(t[0].top_logprobs[0].token, "Hi");
354        assert_eq!(t[0].top_logprobs[1].token, "Hello");
355        let list = r#"{"choices":[{"logprobs":{"tokens":["Hi"],"token_logprobs":[-0.1],
356            "top_logprobs":[[{"token":"Hi","logprob":-0.1},{"token":"Yo","logprob":-3.0}]]}}]}"#;
357        let t = parse_logprobs(list).unwrap();
358        assert_eq!(t[0].top_logprobs.len(), 2);
359        let bad = r#"{"choices":[{"logprobs":{"tokens":["a","b"],"token_logprobs":[-0.1]}}]}"#;
360        assert!(parse_logprobs(bad).is_err());
361    }
362
363    #[test]
364    fn special_tokens_are_dropped() {
365        let t = parse_logprobs(
366            r#"[{"token":"Hi","logprob":-0.1},{"token":"<|eot_id|>","logprob":-3.0}]"#,
367        )
368        .unwrap();
369        assert_eq!(t.len(), 1);
370    }
371
372    #[test]
373    fn missing_or_null_logprobs_is_a_clear_error() {
374        let e = parse_logprobs(r#"{"choices":[{"logprobs":null}]}"#).unwrap_err();
375        assert!(e.to_string().contains("logprobs"));
376        let e = parse_logprobs(r#"{"error":{"message":"bad key"}}"#).unwrap_err();
377        assert!(e.to_string().contains("bad key"));
378    }
379
380    #[test]
381    fn flags_whole_word_when_a_subword_token_is_low() {
382        let t = vec![
383            tok("In", 0.9),
384            tok(" D", 0.95),
385            tok("ord", 0.3),
386            tok("recht", 1.0),
387        ];
388        let r = detect(&t, 0.5);
389        assert_eq!(r.spans.len(), 1);
390        assert_eq!(r.spans[0].text, " Dordrecht");
391        assert_eq!((r.spans[0].start, r.spans[0].end), (1, 4));
392        assert_eq!(r.spans[0].weakest_token, "ord");
393    }
394
395    #[test]
396    fn punctuation_and_whitespace_never_trigger() {
397        let t = vec![
398            tok("Yes", 0.9),
399            tok(",", 0.1),
400            tok(" ", 0.1),
401            tok("ok", 0.9),
402        ];
403        assert!(!detect(&t, 0.5).flagged());
404    }
405
406    #[test]
407    fn adjacent_low_words_merge_into_one_span() {
408        let t = vec![
409            tok("It", 0.99),
410            tok(" was", 0.2),
411            tok(" Pete", 0.3),
412            tok(".", 0.99),
413            tok(" Then", 0.1),
414        ];
415        let r = detect(&t, 0.5);
416        assert_eq!(r.spans.len(), 2);
417        assert_eq!(r.spans[0].text, " was Pete");
418        assert_eq!(r.spans[1].text, " Then");
419    }
420
421    #[test]
422    fn threshold_is_strict_less_than() {
423        let t = vec![tok("a", 0.5)];
424        assert!(!detect(&t, 0.5).flagged());
425        assert!(detect(&t, 0.51).flagged());
426    }
427
428    #[test]
429    fn digits_split_by_space_token_start_a_new_word() {
430        let t = vec![
431            tok(" in", 0.99),
432            tok(" ", 0.99),
433            tok("169", 0.2),
434            tok("1", 0.99),
435        ];
436        let r = detect(&t, 0.5);
437        assert_eq!(r.spans.len(), 1);
438        assert_eq!(r.spans[0].text, "1691");
439    }
440
441    #[test]
442    fn describe_lists_alternatives_but_not_the_chosen_token() {
443        let mut w = tok("ord", 0.57);
444        w.top_logprobs = vec![
445            Alternative {
446                token: "ord".into(),
447                logprob: 0.57f64.ln(),
448            },
449            Alternative {
450                token: "üsseldorf".into(),
451                logprob: 0.39f64.ln(),
452            },
453        ];
454        let r = detect(&[tok(" D", 0.99), w], 0.6);
455        let d = r.spans[0].describe();
456        assert!(d.contains("p=0.57"), "{d}");
457        assert!(d.contains("üsseldorf"), "{d}");
458        assert_eq!(d.matches("\"ord\"").count(), 1, "{d}");
459    }
460
461    #[test]
462    fn bundled_sample_flags_the_birthplace() {
463        let json = include_str!("../examples/logprobs/cuyp.json");
464        let tokens = parse_logprobs(json).unwrap();
465        let r = detect(&tokens, 0.6);
466        assert_eq!(
467            r.text,
468            "Aelbert Cuyp died in 1691 in Dordrecht, Netherlands."
469        );
470        assert!(r.spans.iter().any(|s| s.text == " Dordrecht"));
471        let a = to_token_analysis(&tokens, &r);
472        assert!(a.validate().is_ok());
473    }
474}