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 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 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
124fn 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
137fn 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
150pub 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 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 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 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
240pub 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}