1use anyhow::{bail, Context, Result};
14use serde::{Deserialize, Serialize};
15use serde_json::Value;
16
17use crate::data::{TokenAnalysis, TokenFlag, TokenInfo};
18
19pub const DEFAULT_THRESHOLD: f64 = 0.5;
21
22pub 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#[derive(Debug, Clone, Serialize, PartialEq)]
47pub struct Span {
48 pub start: usize,
50 pub end: usize,
52 pub text: String,
53 pub weakest_token: String,
55 pub min_prob: f64,
56 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
82pub 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 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 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
130fn 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
182fn 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
195fn 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
208pub 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 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 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 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
298pub 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}