1use std::collections::HashMap;
19
20#[derive(Debug, Clone, PartialEq, Eq, Hash)]
24pub enum TokenizerFamily {
25 GPT4,
27 Claude,
29 Gemini,
31 Llama,
33 Generic,
35}
36
37#[derive(Debug, Clone, PartialEq, Eq)]
39pub enum CountMethod {
40 Exact,
42 BpeApprox,
44 CharacterBased,
46 WordBased,
48}
49
50#[derive(Debug, Clone)]
54pub struct TokenCount {
55 pub input_tokens: usize,
57 pub output_tokens: usize,
59 pub total_tokens: usize,
61 pub model: String,
63 pub method: CountMethod,
65}
66
67#[derive(Debug, Clone)]
69pub struct BpeApproxTokenizer {
70 pub family: TokenizerFamily,
72}
73
74impl BpeApproxTokenizer {
75 pub fn count_tokens(text: &str, family: &TokenizerFamily) -> usize {
83 if text.is_empty() {
84 return 0;
85 }
86
87 let chars_per_token: f64 = match family {
88 TokenizerFamily::GPT4 => 4.0,
89 TokenizerFamily::Claude => 3.8,
90 TokenizerFamily::Gemini => 3.9,
91 TokenizerFamily::Llama => 3.6,
92 TokenizerFamily::Generic => 4.0,
93 };
94
95 let char_count = text.chars().count() as f64;
96 let mut estimate = char_count / chars_per_token;
97
98 let punct_count = text
100 .chars()
101 .filter(|c| c.is_ascii_punctuation())
102 .count() as f64;
103 estimate += punct_count * 0.15;
104
105 let digit_count = text.chars().filter(|c| c.is_ascii_digit()).count() as f64;
107 estimate += digit_count * 0.05;
108
109 let newline_count = text.chars().filter(|&c| c == '\n' || c == '\r').count() as f64;
111 estimate += newline_count * 0.3;
112
113 (estimate.ceil() as usize).max(1)
115 }
116
117 pub fn estimate_output_tokens(input_tokens: usize, task_hint: &str) -> usize {
121 let ratio: f64 = {
122 let hint = task_hint.to_lowercase();
123 if hint.contains("summarize") || hint.contains("summarise") {
124 0.25
125 } else if hint.contains("translate") {
126 1.05
127 } else if hint.contains("explain") {
128 1.5
129 } else if hint.contains("code") || hint.contains("implement") || hint.contains("write") {
130 2.0
131 } else if hint.contains("classify") || hint.contains("sentiment") {
132 0.1
133 } else {
134 1.0
135 }
136 };
137 ((input_tokens as f64 * ratio).ceil() as usize).max(1)
138 }
139
140 pub fn split_into_chunks(
146 text: &str,
147 max_tokens: usize,
148 overlap: usize,
149 family: &TokenizerFamily,
150 ) -> Vec<String> {
151 if text.is_empty() || max_tokens == 0 {
152 return vec![];
153 }
154
155 let sentences: Vec<&str> = split_sentences(text);
157
158 let mut chunks: Vec<String> = Vec::new();
159 let mut current = String::new();
160 let mut current_tokens = 0usize;
161
162 for sentence in &sentences {
163 let s_tokens = Self::count_tokens(sentence, family);
164
165 if current_tokens + s_tokens > max_tokens && !current.is_empty() {
166 chunks.push(current.clone());
167 let overlap_buf = build_overlap(¤t, overlap, family);
169 current = overlap_buf.join(" ");
170 current_tokens = Self::count_tokens(¤t, family);
171 }
172
173 if !current.is_empty() {
174 current.push(' ');
175 }
176 current.push_str(sentence);
177 current_tokens += s_tokens;
178 }
179
180 if !current.is_empty() {
181 chunks.push(current);
182 }
183
184 if chunks.is_empty() {
185 chunks.push(text.to_string());
186 }
187
188 chunks
189 }
190
191 pub fn fits_in_context(text: &str, model: &str, context_window: usize, family: &TokenizerFamily) -> bool {
193 let available = (context_window as f64 * 0.95) as usize;
195 let count = Self::count_tokens(text, family);
196 let _ = model; count <= available
198 }
199}
200
201fn split_sentences(text: &str) -> Vec<&str> {
203 let mut sentences: Vec<&str> = Vec::new();
204 let mut start = 0usize;
205 let bytes = text.as_bytes();
206 let len = bytes.len();
207
208 let mut i = 0usize;
209 while i < len {
210 let b = bytes[i];
211 if (b == b'.' || b == b'!' || b == b'?') && i + 1 < len && bytes[i + 1] == b' ' {
212 sentences.push(text[start..=i].trim());
213 start = i + 2;
214 i += 2;
215 } else {
216 i += 1;
217 }
218 }
219 let tail = text[start..].trim();
220 if !tail.is_empty() {
221 sentences.push(tail);
222 }
223 sentences.into_iter().filter(|s| !s.is_empty()).collect()
224}
225
226fn build_overlap(chunk: &str, overlap_tokens: usize, family: &TokenizerFamily) -> Vec<String> {
228 if overlap_tokens == 0 {
229 return vec![];
230 }
231 let words: Vec<&str> = chunk.split_whitespace().collect();
232 let mut buf: Vec<String> = Vec::new();
233 let mut tok_count = 0usize;
234 for word in words.iter().rev() {
235 let wt = BpeApproxTokenizer::count_tokens(word, family);
236 if tok_count + wt > overlap_tokens {
237 break;
238 }
239 buf.push(word.to_string());
240 tok_count += wt;
241 }
242 buf.reverse();
243 buf
244}
245
246pub struct TokenCounter {
250 pub tokenizer: BpeApproxTokenizer,
252 pub model_contexts: HashMap<String, usize>,
254}
255
256impl TokenCounter {
257 pub fn new(family: TokenizerFamily) -> Self {
260 Self {
261 tokenizer: BpeApproxTokenizer { family },
262 model_contexts: Self::build_context_table(),
263 }
264 }
265
266 pub fn count(&self, text: &str) -> TokenCount {
268 let n = BpeApproxTokenizer::count_tokens(text, &self.tokenizer.family);
269 TokenCount {
270 input_tokens: n,
271 output_tokens: 0,
272 total_tokens: n,
273 model: String::new(),
274 method: CountMethod::BpeApprox,
275 }
276 }
277
278 pub fn count_messages(&self, messages: &[(String, String)]) -> TokenCount {
283 const ROLE_OVERHEAD: usize = 4;
284 let n: usize = messages
285 .iter()
286 .map(|(_, content)| {
287 BpeApproxTokenizer::count_tokens(content, &self.tokenizer.family) + ROLE_OVERHEAD
288 })
289 .sum();
290 let total = n + 2;
292 TokenCount {
293 input_tokens: total,
294 output_tokens: 0,
295 total_tokens: total,
296 model: String::new(),
297 method: CountMethod::BpeApprox,
298 }
299 }
300
301 pub fn count_with_system(&self, system: &str, messages: &[(String, String)]) -> TokenCount {
303 let sys_tokens = BpeApproxTokenizer::count_tokens(system, &self.tokenizer.family);
304 let sys_total = sys_tokens + 5;
306 let mut msg_count = self.count_messages(messages);
307 msg_count.input_tokens += sys_total;
308 msg_count.total_tokens += sys_total;
309 msg_count
310 }
311
312 pub fn remaining_context(&self, model: &str, used: usize) -> Option<usize> {
316 let window = self.model_contexts.get(model).copied()
317 .or_else(|| Self::model_context_window(model))?;
318 Some(window.saturating_sub(used))
319 }
320
321 pub fn batch_count(&self, texts: &[&str]) -> Vec<TokenCount> {
323 texts.iter().map(|t| self.count(t)).collect()
324 }
325
326 pub fn model_context_window(model: &str) -> Option<usize> {
330 let table = Self::build_context_table();
331 if let Some(&w) = table.get(model) {
333 return Some(w);
334 }
335 table
337 .iter()
338 .filter(|(k, _)| model.starts_with(k.as_str()))
339 .max_by_key(|(k, _)| k.len())
340 .map(|(_, &v)| v)
341 }
342
343 fn build_context_table() -> HashMap<String, usize> {
346 let entries: &[(&str, usize)] = &[
347 ("gpt-3.5-turbo", 16_385),
349 ("gpt-3.5-turbo-16k", 16_385),
350 ("gpt-4", 8_192),
351 ("gpt-4-32k", 32_768),
352 ("gpt-4-turbo", 128_000),
353 ("gpt-4-turbo-preview", 128_000),
354 ("gpt-4o", 128_000),
355 ("gpt-4o-mini", 128_000),
356 ("gpt-4.5", 128_000),
357 ("o1", 200_000),
358 ("o1-mini", 128_000),
359 ("o3", 200_000),
360 ("o3-mini", 200_000),
361 ("claude-2", 100_000),
363 ("claude-2.1", 200_000),
364 ("claude-3-haiku", 200_000),
365 ("claude-3-sonnet", 200_000),
366 ("claude-3-opus", 200_000),
367 ("claude-3-5-haiku", 200_000),
368 ("claude-3-5-sonnet", 200_000),
369 ("claude-3-5-opus", 200_000),
370 ("claude-3-7-sonnet", 200_000),
371 ("claude-sonnet-4", 200_000),
372 ("claude-opus-4", 200_000),
373 ("gemini-pro", 32_768),
375 ("gemini-1.0-pro", 32_768),
376 ("gemini-1.5-pro", 1_048_576),
377 ("gemini-1.5-flash", 1_048_576),
378 ("gemini-2.0-flash", 1_048_576),
379 ("gemini-2.0-pro", 2_097_152),
380 ("llama-2-7b", 4_096),
382 ("llama-2-13b", 4_096),
383 ("llama-2-70b", 4_096),
384 ("llama-3-8b", 8_192),
385 ("llama-3-70b", 8_192),
386 ("llama-3.1-8b", 131_072),
387 ("llama-3.1-70b", 131_072),
388 ("llama-3.1-405b", 131_072),
389 ("llama-3.3-70b", 131_072),
390 ("mistral-7b", 32_768),
392 ("mistral-8x7b", 32_768),
393 ("mistral-large", 131_072),
394 ("mistral-small", 131_072),
395 ("command-r", 128_000),
397 ("command-r-plus", 128_000),
398 ("deepseek-chat", 64_000),
400 ("deepseek-coder", 16_000),
401 ("deepseek-r1", 163_840),
402 ];
403 entries
404 .iter()
405 .map(|&(k, v)| (k.to_string(), v))
406 .collect()
407 }
408}
409
410#[cfg(test)]
411mod tests {
412 use super::*;
413
414 #[test]
415 fn count_non_zero_for_non_empty() {
416 for family in [
417 TokenizerFamily::GPT4,
418 TokenizerFamily::Claude,
419 TokenizerFamily::Gemini,
420 TokenizerFamily::Llama,
421 TokenizerFamily::Generic,
422 ] {
423 let n = BpeApproxTokenizer::count_tokens("Hello, world!", &family);
424 assert!(n > 0, "family {family:?} returned 0 tokens");
425 }
426 }
427
428 #[test]
429 fn count_empty_is_zero() {
430 assert_eq!(BpeApproxTokenizer::count_tokens("", &TokenizerFamily::GPT4), 0);
431 }
432
433 #[test]
434 fn output_estimate_summarize_less_than_explain() {
435 let summarize = BpeApproxTokenizer::estimate_output_tokens(100, "summarize this");
436 let explain = BpeApproxTokenizer::estimate_output_tokens(100, "explain this concept");
437 assert!(summarize < explain);
438 }
439
440 #[test]
441 fn split_into_chunks_respects_max() {
442 let text = "The quick brown fox. The lazy dog jumped. Over the fence. Back again.";
443 let chunks = BpeApproxTokenizer::split_into_chunks(text, 5, 0, &TokenizerFamily::GPT4);
444 for chunk in &chunks {
445 let t = BpeApproxTokenizer::count_tokens(chunk, &TokenizerFamily::GPT4);
446 assert!(t <= 20, "chunk too large: {t} tokens");
448 }
449 assert!(!chunks.is_empty());
450 }
451
452 #[test]
453 fn context_window_known_model() {
454 assert_eq!(
455 TokenCounter::model_context_window("gpt-4o"),
456 Some(128_000)
457 );
458 assert_eq!(
459 TokenCounter::model_context_window("claude-3-5-sonnet"),
460 Some(200_000)
461 );
462 }
463
464 #[test]
465 fn context_window_unknown_model() {
466 assert_eq!(
467 TokenCounter::model_context_window("totally-made-up-model-9999"),
468 None
469 );
470 }
471
472 #[test]
473 fn count_messages_adds_overhead() {
474 let counter = TokenCounter::new(TokenizerFamily::GPT4);
475 let msgs = vec![
476 ("user".to_string(), "Hi".to_string()),
477 ("assistant".to_string(), "Hello!".to_string()),
478 ];
479 let single_text = counter.count("HiHello!");
480 let msg_count = counter.count_messages(&msgs);
481 assert!(msg_count.total_tokens > single_text.total_tokens);
483 }
484
485 #[test]
486 fn remaining_context() {
487 let counter = TokenCounter::new(TokenizerFamily::GPT4);
488 let rem = counter.remaining_context("gpt-4o", 1000);
489 assert_eq!(rem, Some(128_000 - 1000));
490 assert_eq!(counter.remaining_context("unknown-model", 0), None);
491 }
492
493 #[test]
494 fn batch_count_length_matches() {
495 let counter = TokenCounter::new(TokenizerFamily::Claude);
496 let texts = ["hello", "world", "foo bar baz"];
497 let results = counter.batch_count(&texts);
498 assert_eq!(results.len(), 3);
499 }
500}