1use std::sync::Arc;
13
14use tiktoken_rs::CoreBPE;
15use tiktoken_rs::cl100k_base;
16
17pub trait Tokenizer: Send + Sync {
19 fn count(&self, text: &str) -> u32;
21 fn name(&self) -> &str;
23}
24
25pub struct TiktokenTokenizer {
27 bpe: Arc<CoreBPE>,
28 name: &'static str,
29}
30
31impl TiktokenTokenizer {
32 pub fn cl100k() -> Result<Self, crate::ClassifiedError> {
39 let bpe = cl100k_base()
40 .map_err(|e| crate::ClassifiedError::Runtime(format!("tiktoken cl100k_base: {e}")))?;
41 Ok(Self {
42 bpe: Arc::new(bpe),
43 name: "tiktoken.cl100k",
44 })
45 }
46}
47
48impl Tokenizer for TiktokenTokenizer {
49 fn count(&self, text: &str) -> u32 {
50 self.bpe
51 .encode_with_special_tokens(text)
52 .len()
53 .try_into()
54 .unwrap_or(u32::MAX / 2)
55 }
56
57 fn name(&self) -> &str {
58 self.name
59 }
60}
61
62#[derive(Debug, Default, Clone, Copy)]
68pub struct ApproxTokenizer;
69
70impl Tokenizer for ApproxTokenizer {
71 fn count(&self, text: &str) -> u32 {
72 let chars = text.chars().count();
73 (chars as u32).div_ceil(4)
74 }
75
76 fn name(&self) -> &str {
77 "approx.chars_over_4"
78 }
79}
80
81#[cfg(test)]
82#[allow(clippy::expect_used, clippy::unwrap_used)]
83mod tests {
84 use super::*;
85
86 #[test]
87 fn approx_counts_nonzero_for_ascii() {
88 let t = ApproxTokenizer;
89 assert_eq!(t.count(""), 0);
90 assert_eq!(t.count("abcd"), 1);
91 assert_eq!(t.count("abcde"), 2);
92 }
93
94 #[test]
95 fn approx_counts_multibyte_by_char_not_byte() {
96 let t = ApproxTokenizer;
97 assert_eq!(t.count("中"), 1);
99 assert_eq!(t.count("中国加油"), 1);
101 }
102
103 #[test]
104 fn tiktoken_counts_standard_phrase() {
105 let t = TiktokenTokenizer::cl100k().unwrap();
106 let n = t.count("Hello world");
108 assert!(n > 0 && n < 10, "expected a small token count, got {n}");
109 }
110
111 #[test]
112 fn tiktoken_empty_is_zero() {
113 let t = TiktokenTokenizer::cl100k().unwrap();
114 assert_eq!(t.count(""), 0);
115 }
116}