1use anyhow::{bail, Context, Result};
14use serde::{Deserialize, Serialize};
15use serde_json::Value;
16use unicode_segmentation::UnicodeSegmentation;
17
18use crate::data::{TokenAnalysis, TokenFlag, TokenInfo};
19
20pub const DEFAULT_THRESHOLD: f64 = 0.5;
22
23pub const FLAG_LABEL: &str = "uncertain";
25
26#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
27pub struct Alternative {
29 pub token: String,
31 pub logprob: f64,
33}
34
35#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
36pub struct LogprobToken {
38 pub token: String,
40 pub logprob: f64,
42 #[serde(default)]
43 pub top_logprobs: Vec<Alternative>,
45}
46
47impl LogprobToken {
48 pub fn prob(&self) -> f64 {
50 self.logprob.exp().clamp(0.0, 1.0)
51 }
52}
53
54#[derive(Debug, Clone, Serialize, PartialEq)]
56pub struct Span {
57 pub start: usize,
59 pub end: usize,
61 pub text: String,
63 pub weakest_token: String,
65 pub min_prob: f64,
67 pub alternatives: Vec<Candidate>,
70 pub entropy_bits: Option<f64>,
76}
77
78#[derive(Debug, Clone, Serialize, PartialEq)]
79pub struct Candidate {
81 pub token: String,
83 pub prob: f64,
85}
86
87#[derive(Debug, Clone, Serialize, PartialEq)]
88pub struct Report {
90 pub threshold: f64,
92 pub text: String,
94 pub token_count: usize,
96 pub mean_prob: f64,
98 pub perplexity: f64,
101 pub spans: Vec<Span>,
103}
104
105impl Report {
106 pub fn flagged(&self) -> bool {
108 !self.spans.is_empty()
109 }
110}
111
112pub fn parse_logprobs(json: &str) -> Result<Vec<LogprobToken>> {
116 let value: Value = serde_json::from_str(json).context("input is not valid JSON")?;
117 if let Some(err) = value.get("error") {
118 bail!("the API returned an error: {}", err);
119 }
120 let legacy = value
123 .pointer("/choices/0/logprobs")
124 .or_else(|| Some(&value).filter(|v| v.get("token_logprobs").is_some()))
125 .filter(|l| l.get("tokens").is_some() && l.get("token_logprobs").is_some());
126 if let Some(l) = legacy {
127 let mut tokens = parse_legacy(l)?;
128 tokens.retain(|t| !is_special_token(&t.token));
129 if tokens.is_empty() {
130 bail!("logprobs contain no tokens");
131 }
132 return Ok(tokens);
133 }
134 if let Some(tokens) = parse_gemini(&value)? {
135 return Ok(tokens);
136 }
137 let content = if value.is_array() {
138 &value
139 } else if let Some(c) = value.pointer("/choices/0/logprobs/content") {
140 c
141 } else if let Some(c) = value.get("content") {
142 c
143 } else {
144 bail!(
145 "no token logprobs found; expected choices[0].logprobs.content \
146 (request the completion with \"logprobs\": true)"
147 );
148 };
149 if content.is_null() {
150 bail!("logprobs are null; request the completion with \"logprobs\": true");
151 }
152 let mut tokens: Vec<LogprobToken> =
155 Vec::<LogprobToken>::deserialize(content).context("could not read logprobs tokens")?;
156 tokens.retain(|t| !is_special_token(&t.token));
159 if tokens.is_empty() {
160 bail!("logprobs contain no tokens");
161 }
162 Ok(tokens)
163}
164
165fn parse_legacy(l: &Value) -> Result<Vec<LogprobToken>> {
169 let toks = l["tokens"]
170 .as_array()
171 .context("logprobs.tokens is not a list")?;
172 let lps = l["token_logprobs"]
173 .as_array()
174 .context("logprobs.token_logprobs is not a list")?;
175 if toks.len() != lps.len() {
176 bail!("logprobs.tokens and logprobs.token_logprobs have different lengths");
177 }
178 let tops = l.get("top_logprobs").and_then(Value::as_array);
179 let mut out = Vec::with_capacity(toks.len());
180 for (i, (t, lp)) in toks.iter().zip(lps).enumerate() {
181 let token = t
182 .as_str()
183 .context("a logprobs token is not a string")?
184 .to_string();
185 let logprob = lp.as_f64().unwrap_or(f64::NEG_INFINITY);
186 let mut top_logprobs: Vec<Alternative> = match tops.and_then(|a| a.get(i)) {
187 Some(Value::Object(m)) => m
188 .iter()
189 .filter_map(|(k, v)| {
190 v.as_f64().map(|lp| Alternative {
191 token: k.clone(),
192 logprob: lp,
193 })
194 })
195 .collect(),
196 Some(v @ Value::Array(_)) => serde_json::from_value(v.clone()).unwrap_or_default(),
197 _ => Vec::new(),
198 };
199 top_logprobs.sort_by(|a, b| b.logprob.total_cmp(&a.logprob));
200 out.push(LogprobToken {
201 token,
202 logprob,
203 top_logprobs,
204 });
205 }
206 Ok(out)
207}
208
209fn parse_gemini(value: &Value) -> Result<Option<Vec<LogprobToken>>> {
214 let lr = value
215 .pointer("/candidates/0/logprobsResult")
216 .or_else(|| value.get("logprobsResult"));
217 let Some(lr) = lr else {
218 return Ok(None);
219 };
220 let chosen = lr
221 .get("chosenCandidates")
222 .and_then(Value::as_array)
223 .context("logprobsResult has no chosenCandidates list")?;
224 let tops = lr.get("topCandidates").and_then(Value::as_array);
225 let cand = |c: &Value| -> Option<Alternative> {
226 Some(Alternative {
227 token: c.get("token")?.as_str()?.to_string(),
228 logprob: c.get("logProbability")?.as_f64()?,
229 })
230 };
231 let mut out = Vec::with_capacity(chosen.len());
232 for (i, c) in chosen.iter().enumerate() {
233 let Some(a) = cand(c) else {
234 bail!("chosenCandidates[{i}] needs a token and a logProbability");
235 };
236 let mut top: Vec<Alternative> = tops
237 .and_then(|t| t.get(i))
238 .and_then(|t| t.get("candidates"))
239 .and_then(Value::as_array)
240 .map(|v| v.iter().filter_map(cand).collect())
241 .unwrap_or_default();
242 top.sort_by(|a, b| b.logprob.total_cmp(&a.logprob));
243 out.push(LogprobToken {
244 token: a.token,
245 logprob: a.logprob,
246 top_logprobs: top,
247 });
248 }
249 out.retain(|t| !is_special_token(&t.token));
250 if out.is_empty() {
251 bail!("logprobs contain no tokens");
252 }
253 Ok(Some(out))
254}
255
256fn is_special_token(s: &str) -> bool {
257 s.starts_with("<|") && s.ends_with("|>")
258}
259
260fn has_word_chars(s: &str) -> bool {
261 s.chars().any(|c| c.is_alphanumeric())
262}
263
264fn words(tokens: &[LogprobToken]) -> Vec<(usize, usize)> {
273 let text: String = tokens.iter().map(|t| t.token.as_str()).collect();
274 let mut boundaries = text.split_word_bound_indices().map(|(i, _)| i).peekable();
277 let mut out: Vec<(usize, usize)> = Vec::new();
278 let mut offset = 0usize;
279 for (i, t) in tokens.iter().enumerate() {
280 while boundaries.next_if(|&b| b < offset).is_some() {}
281 let on_boundary = boundaries.peek() == Some(&offset);
282 if out.is_empty() || on_boundary {
283 out.push((i, i + 1));
284 } else if let Some(last) = out.last_mut() {
285 last.1 = i + 1;
286 }
287 offset += t.token.len();
288 }
289 out
290}
291
292pub fn detect(tokens: &[LogprobToken], threshold: f64) -> Report {
294 let text: String = tokens.iter().map(|t| t.token.as_str()).collect();
295 let mean_prob = if tokens.is_empty() {
296 0.0
297 } else {
298 tokens.iter().map(|t| t.prob()).sum::<f64>() / tokens.len() as f64
299 };
300 let perplexity = if tokens.is_empty() {
301 1.0
302 } else {
303 let mean_lp = tokens.iter().map(|t| t.logprob).sum::<f64>() / tokens.len() as f64;
304 (-mean_lp).exp()
305 };
306
307 let flagged_words: Vec<(usize, usize)> = words(tokens)
310 .into_iter()
311 .filter(|&(s, e)| {
312 tokens[s..e]
313 .iter()
314 .any(|t| has_word_chars(&t.token) && t.prob() < threshold)
315 })
316 .collect();
317
318 let mut ranges: Vec<(usize, usize)> = Vec::new();
320 for (s, e) in flagged_words {
321 if let Some(last) = ranges.last_mut() {
322 if tokens[last.1..s].iter().all(|t| t.token.trim().is_empty()) {
323 last.1 = e;
324 continue;
325 }
326 }
327 ranges.push((s, e));
328 }
329
330 let spans = ranges
331 .into_iter()
332 .map(|(start, end)| {
333 let weakest = (start..end)
334 .filter(|&i| has_word_chars(&tokens[i].token))
335 .min_by(|&a, &b| tokens[a].prob().total_cmp(&tokens[b].prob()))
336 .unwrap_or(start);
337 let w = &tokens[weakest];
338 let alternatives: Vec<Candidate> = w
339 .top_logprobs
340 .iter()
341 .map(|a| Candidate {
342 token: a.token.clone(),
343 prob: a.logprob.exp(),
344 })
345 .collect();
346 let entropy_bits = top_k_entropy_bits(&alternatives);
347 Span {
348 start,
349 end,
350 text: tokens[start..end]
351 .iter()
352 .map(|t| t.token.as_str())
353 .collect(),
354 weakest_token: w.token.clone(),
355 min_prob: w.prob(),
356 alternatives,
357 entropy_bits,
358 }
359 })
360 .collect();
361
362 Report {
363 threshold,
364 text,
365 token_count: tokens.len(),
366 mean_prob,
367 perplexity,
368 spans,
369 }
370}
371
372pub fn top_k_entropy_bits(alternatives: &[Candidate]) -> Option<f64> {
375 let total: f64 = alternatives.iter().map(|c| c.prob).filter(|p| p.is_finite() && *p > 0.0).sum();
376 if alternatives.is_empty() || total <= 0.0 {
377 return None;
378 }
379 Some(
380 alternatives
381 .iter()
382 .map(|c| c.prob / total)
383 .filter(|p| p.is_finite() && *p > 0.0)
384 .map(|p| -p * p.log2())
385 .sum::<f64>()
386 .max(0.0),
387 )
388}
389
390impl Span {
391 pub fn describe(&self) -> String {
394 let mut s = format!("p={:.2} at {:?}", self.min_prob, self.weakest_token);
395 let others: Vec<String> = self
396 .alternatives
397 .iter()
398 .filter(|c| c.token != self.weakest_token)
399 .map(|c| format!("{:?} ({:.2})", c.token, c.prob))
400 .collect();
401 if !others.is_empty() {
402 s.push_str("; model also considered ");
403 s.push_str(&others.join(", "));
404 }
405 s
406 }
407}
408
409pub fn to_token_analysis(tokens: &[LogprobToken], report: &Report) -> TokenAnalysis {
412 TokenAnalysis {
413 tokens: tokens
414 .iter()
415 .map(|t| TokenInfo {
416 text: t.token.clone(),
417 confidence: t.prob(),
418 })
419 .collect(),
420 flags: report
421 .spans
422 .iter()
423 .map(|s| TokenFlag {
424 start: s.start,
425 end: s.end,
426 flag: FLAG_LABEL.to_string(),
427 description: Some(s.describe()),
428 })
429 .collect(),
430 }
431}
432
433#[cfg(test)]
434mod tests {
435 use super::*;
436
437 fn tok(t: &str, p: f64) -> LogprobToken {
438 LogprobToken {
439 token: t.to_string(),
440 logprob: p.ln(),
441 top_logprobs: vec![],
442 }
443 }
444
445 #[test]
446 fn parses_all_three_shapes() {
447 let arr = r#"[{"token":"Hi","logprob":-0.1}]"#;
448 let obj = r#"{"content":[{"token":"Hi","logprob":-0.1}]}"#;
449 let full = r#"{"choices":[{"logprobs":{"content":[{"token":"Hi","logprob":-0.1,"top_logprobs":[]}]}}]}"#;
450 for s in [arr, obj, full] {
451 let t = parse_logprobs(s).unwrap();
452 assert_eq!(t.len(), 1);
453 assert_eq!(t[0].token, "Hi");
454 }
455 }
456
457 #[test]
458 fn parses_tokens_and_token_logprobs_shape() {
459 let map = r#"{"choices":[{"logprobs":{"tokens":["Hi","!"],"token_logprobs":[-0.1,-2.0],
460 "top_logprobs":[{"Hello":-2.5,"Hi":-0.1},{"!":-2.0}]}}]}"#;
461 let t = parse_logprobs(map).unwrap();
462 assert_eq!(t.len(), 2);
463 assert_eq!(t[0].token, "Hi");
464 assert_eq!(t[0].top_logprobs[0].token, "Hi");
465 assert_eq!(t[0].top_logprobs[1].token, "Hello");
466 let list = r#"{"choices":[{"logprobs":{"tokens":["Hi"],"token_logprobs":[-0.1],
467 "top_logprobs":[[{"token":"Hi","logprob":-0.1},{"token":"Yo","logprob":-3.0}]]}}]}"#;
468 let t = parse_logprobs(list).unwrap();
469 assert_eq!(t[0].top_logprobs.len(), 2);
470 let bad = r#"{"choices":[{"logprobs":{"tokens":["a","b"],"token_logprobs":[-0.1]}}]}"#;
471 assert!(parse_logprobs(bad).is_err());
472 }
473
474 #[test]
475 fn special_tokens_are_dropped() {
476 let t = parse_logprobs(
477 r#"[{"token":"Hi","logprob":-0.1},{"token":"<|eot_id|>","logprob":-3.0}]"#,
478 )
479 .unwrap();
480 assert_eq!(t.len(), 1);
481 }
482
483 #[test]
484 fn missing_or_null_logprobs_is_a_clear_error() {
485 let e = parse_logprobs(r#"{"choices":[{"logprobs":null}]}"#).unwrap_err();
486 assert!(e.to_string().contains("logprobs"));
487 let e = parse_logprobs(r#"{"error":{"message":"bad key"}}"#).unwrap_err();
488 assert!(e.to_string().contains("bad key"));
489 }
490
491 #[test]
492 fn flags_whole_word_when_a_subword_token_is_low() {
493 let t = vec![
494 tok("In", 0.9),
495 tok(" D", 0.95),
496 tok("ord", 0.3),
497 tok("recht", 1.0),
498 ];
499 let r = detect(&t, 0.5);
500 assert_eq!(r.spans.len(), 1);
501 assert_eq!(r.spans[0].text, " Dordrecht");
502 assert_eq!((r.spans[0].start, r.spans[0].end), (1, 4));
503 assert_eq!(r.spans[0].weakest_token, "ord");
504 }
505
506 #[test]
507 fn punctuation_and_whitespace_never_trigger() {
508 let t = vec![
509 tok("Yes", 0.9),
510 tok(",", 0.1),
511 tok(" ", 0.1),
512 tok("ok", 0.9),
513 ];
514 assert!(!detect(&t, 0.5).flagged());
515 }
516
517 #[test]
518 fn adjacent_low_words_merge_into_one_span() {
519 let t = vec![
520 tok("It", 0.99),
521 tok(" was", 0.2),
522 tok(" Pete", 0.3),
523 tok(".", 0.99),
524 tok(" Then", 0.1),
525 ];
526 let r = detect(&t, 0.5);
527 assert_eq!(r.spans.len(), 2);
528 assert_eq!(r.spans[0].text, " was Pete");
529 assert_eq!(r.spans[1].text, " Then");
530 }
531
532 #[test]
533 fn threshold_is_strict_less_than() {
534 let t = vec![tok("a", 0.5)];
535 assert!(!detect(&t, 0.5).flagged());
536 assert!(detect(&t, 0.51).flagged());
537 }
538
539 #[test]
540 fn digits_split_by_space_token_start_a_new_word() {
541 let t = vec![
542 tok(" in", 0.99),
543 tok(" ", 0.99),
544 tok("169", 0.2),
545 tok("1", 0.99),
546 ];
547 let r = detect(&t, 0.5);
548 assert_eq!(r.spans.len(), 1);
549 assert_eq!(r.spans[0].text, "1691");
550 }
551
552 #[test]
553 fn describe_lists_alternatives_but_not_the_chosen_token() {
554 let mut w = tok("ord", 0.57);
555 w.top_logprobs = vec![
556 Alternative {
557 token: "ord".into(),
558 logprob: 0.57f64.ln(),
559 },
560 Alternative {
561 token: "üsseldorf".into(),
562 logprob: 0.39f64.ln(),
563 },
564 ];
565 let r = detect(&[tok(" D", 0.99), w], 0.6);
566 let d = r.spans[0].describe();
567 assert!(d.contains("p=0.57"), "{d}");
568 assert!(d.contains("üsseldorf"), "{d}");
569 assert_eq!(d.matches("\"ord\"").count(), 1, "{d}");
570 }
571
572 #[test]
573 fn chinese_answer_flags_the_unsure_character_not_the_sentence() {
574 let t = vec![
576 tok("埃", 0.99),
577 tok("菲", 0.99),
578 tok("尔", 0.99),
579 tok("铁", 0.99),
580 tok("塔", 0.99),
581 tok("在", 0.99),
582 tok("巴", 0.95),
583 tok("黎", 0.30),
584 tok("。", 0.99),
585 ];
586 let r = detect(&t, 0.5);
587 assert_eq!(r.spans.len(), 1);
588 assert_eq!(r.spans[0].text, "黎");
590 }
591
592 #[test]
593 fn contractions_and_hyphens_follow_unicode_rules() {
594 let t = vec![tok(" don", 0.9), tok("'t", 0.2), tok(" well", 0.9), tok("-", 0.9), tok("known", 0.9)];
595 let r = detect(&t, 0.5);
596 assert_eq!(r.spans.len(), 1);
597 assert_eq!(r.spans[0].text, " don't");
598 }
599
600 #[test]
601 fn perplexity_and_entropy() {
602 let t = vec![tok("a", 1.0), tok("b", 1.0)];
603 assert!((detect(&t, 0.5).perplexity - 1.0).abs() < 1e-12);
604 let t = vec![tok("a", 0.5), tok("b", 0.5)];
605 assert!((detect(&t, 0.5).perplexity - 2.0).abs() < 1e-9);
606 let two = [
607 Candidate { token: "x".into(), prob: 0.4 },
608 Candidate { token: "y".into(), prob: 0.4 },
609 ];
610 assert!((top_k_entropy_bits(&two).unwrap() - 1.0).abs() < 1e-12);
611 assert_eq!(top_k_entropy_bits(&[]), None);
612 }
613
614 #[test]
615 fn parses_gemini_logprobs_result() {
616 let json = r#"{"candidates":[{"content":{"parts":[{"text":"Paris."}]},
617 "logprobsResult":{
618 "topCandidates":[
619 {"candidates":[{"token":"Paris","logProbability":-0.05},{"token":"Lyon","logProbability":-3.2}]},
620 {"candidates":[{"token":".","logProbability":-0.01}]}],
621 "chosenCandidates":[{"token":"Paris","logProbability":-0.05},{"token":".","logProbability":-0.01}]}}]}"#;
622 let t = parse_logprobs(json).unwrap();
623 assert_eq!(t.len(), 2);
624 assert_eq!(t[0].token, "Paris");
625 assert_eq!(t[0].top_logprobs[1].token, "Lyon");
626 assert!(parse_logprobs(r#"{"logprobsResult":{"topCandidates":[]}}"#).is_err());
627 }
628
629 #[test]
630 fn bundled_sample_flags_the_birthplace() {
631 let json = include_str!("../examples/logprobs/cuyp.json");
632 let tokens = parse_logprobs(json).unwrap();
633 let r = detect(&tokens, 0.6);
634 assert_eq!(
635 r.text,
636 "Aelbert Cuyp died in 1691 in Dordrecht, Netherlands."
637 );
638 assert!(r.spans.iter().any(|s| s.text == " Dordrecht"));
639 let a = to_token_analysis(&tokens, &r);
640 assert!(a.validate().is_ok());
641 }
642}