Skip to main content

tokio_prompt_orchestrator/
token_counter.rs

1//! Multi-Model Token Counting with BPE Approximation
2//!
3//! Provides heuristic token counting for all major LLM families without
4//! requiring a full tokeniser dependency. Counts are BPE approximations
5//! tuned per-family and are accurate to within ~5–10 % for typical English
6//! prose and code.
7//!
8//! ## Quick Start
9//!
10//! ```rust
11//! use tokio_prompt_orchestrator::token_counter::{TokenCounter, TokenizerFamily};
12//!
13//! let counter = TokenCounter::new(TokenizerFamily::Claude);
14//! let tc = counter.count("Hello, world! This is a test.");
15//! println!("tokens: {}", tc.total_tokens);
16//! ```
17
18use std::collections::HashMap;
19
20// ── Enums ────────────────────────────────────────────────────────────────────
21
22/// The tokeniser family that governs counting heuristics.
23#[derive(Debug, Clone, PartialEq, Eq, Hash)]
24pub enum TokenizerFamily {
25    /// GPT-3.5 / GPT-4 (cl100k_base) — ~4 chars per token.
26    GPT4,
27    /// Claude 2 / 3 family — ~3.8 chars per token (slightly denser vocab).
28    Claude,
29    /// Google Gemini — ~3.9 chars per token.
30    Gemini,
31    /// Meta LLaMA / LLaMA-2 — ~3.6 chars per token (SentencePiece).
32    Llama,
33    /// Unknown family — falls back to 4 chars per token.
34    Generic,
35}
36
37/// How the token count was derived.
38#[derive(Debug, Clone, PartialEq, Eq)]
39pub enum CountMethod {
40    /// Counted by the model's actual tokeniser (not yet implemented here).
41    Exact,
42    /// Approximated via BPE character-ratio heuristic.
43    BpeApprox,
44    /// Approximated by raw character count.
45    CharacterBased,
46    /// Approximated by whitespace-split word count.
47    WordBased,
48}
49
50// ── Structs ───────────────────────────────────────────────────────────────────
51
52/// Result of a token-count operation.
53#[derive(Debug, Clone)]
54pub struct TokenCount {
55    /// Estimated input tokens.
56    pub input_tokens: usize,
57    /// Estimated output tokens (0 unless a task-hint estimation was used).
58    pub output_tokens: usize,
59    /// `input_tokens + output_tokens`.
60    pub total_tokens: usize,
61    /// Model identifier this count was computed for.
62    pub model: String,
63    /// Method used to derive the count.
64    pub method: CountMethod,
65}
66
67/// BPE-approximation tokeniser parameterised by family.
68#[derive(Debug, Clone)]
69pub struct BpeApproxTokenizer {
70    /// The tokeniser family controlling the chars-per-token ratio.
71    pub family: TokenizerFamily,
72}
73
74impl BpeApproxTokenizer {
75    /// Approximate token count for `text` using family-specific heuristics.
76    ///
77    /// The heuristic:
78    /// 1. Compute a base count from `chars / chars_per_token`.
79    /// 2. Add +1 token per punctuation run (commas, periods, etc.).
80    /// 3. Adjust for runs of digits (numbers tokenise more densely).
81    /// 4. Adjust for leading/trailing whitespace overhead.
82    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        // Punctuation bonus: each standalone punctuation char adds ~0.3 tokens.
99        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        // Digit runs: numbers split into smaller tokens.
106        let digit_count = text.chars().filter(|c| c.is_ascii_digit()).count() as f64;
107        estimate += digit_count * 0.05;
108
109        // Whitespace: newlines and tabs each cost an extra fraction.
110        let newline_count = text.chars().filter(|&c| c == '\n' || c == '\r').count() as f64;
111        estimate += newline_count * 0.3;
112
113        // Always at least 1 token for non-empty text.
114        (estimate.ceil() as usize).max(1)
115    }
116
117    /// Estimate output tokens given input token count and a task hint.
118    ///
119    /// The ratios are empirical defaults; callers should tune for their workload.
120    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    /// Split `text` into chunks of at most `max_tokens` tokens, with `overlap`
141    /// tokens of context carried forward between chunks.
142    ///
143    /// Chunks are split at sentence boundaries (`.`, `!`, `?`) where possible
144    /// to preserve semantic coherence.
145    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        // Split into sentences on common terminators.
156        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                // Build overlap from tail of current chunk.
168                let overlap_buf = build_overlap(&current, overlap, family);
169                current = overlap_buf.join(" ");
170                current_tokens = Self::count_tokens(&current, 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    /// Return `true` if `text` fits within `context_window` tokens for `model`.
192    pub fn fits_in_context(text: &str, model: &str, context_window: usize, family: &TokenizerFamily) -> bool {
193        // Add ~5 % overhead for system prompts and role tokens.
194        let available = (context_window as f64 * 0.95) as usize;
195        let count = Self::count_tokens(text, family);
196        let _ = model; // model name kept for API symmetry / future lookup
197        count <= available
198    }
199}
200
201/// Split text into sentences on `.`, `!`, `?` followed by whitespace.
202fn 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
226/// Build an overlap buffer from the tail of `chunk`.
227fn 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
246// ── TokenCounter ──────────────────────────────────────────────────────────────
247
248/// High-level token counter for a specific model/family combination.
249pub struct TokenCounter {
250    /// Underlying BPE approximation tokeniser.
251    pub tokenizer: BpeApproxTokenizer,
252    /// Per-model context window sizes (tokens).
253    pub model_contexts: HashMap<String, usize>,
254}
255
256impl TokenCounter {
257    /// Create a counter for the given tokeniser family, pre-populated with
258    /// context windows for all major models.
259    pub fn new(family: TokenizerFamily) -> Self {
260        Self {
261            tokenizer: BpeApproxTokenizer { family },
262            model_contexts: Self::build_context_table(),
263        }
264    }
265
266    /// Count tokens in a single text string.
267    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    /// Count tokens across a list of `(role, content)` message pairs.
279    ///
280    /// Each message incurs a ~4-token overhead for the role marker and
281    /// formatting tokens used by chat APIs.
282    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        // Add 2 tokens for the reply-primer used by most chat completions APIs.
291        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    /// Count tokens for a system prompt plus a message list.
302    pub fn count_with_system(&self, system: &str, messages: &[(String, String)]) -> TokenCount {
303        let sys_tokens = BpeApproxTokenizer::count_tokens(system, &self.tokenizer.family);
304        // System prompts have ~5 token framing overhead.
305        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    /// How many tokens remain in `model`'s context window after `used` tokens.
313    ///
314    /// Returns `None` if the model is not in the built-in table.
315    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    /// Count tokens for each text in `texts`, returning one [`TokenCount`] per entry.
322    pub fn batch_count(&self, texts: &[&str]) -> Vec<TokenCount> {
323        texts.iter().map(|t| self.count(t)).collect()
324    }
325
326    /// Look up the context window (in tokens) for a known model.
327    ///
328    /// Returns `None` for unrecognised model identifiers.
329    pub fn model_context_window(model: &str) -> Option<usize> {
330        let table = Self::build_context_table();
331        // Also try prefix matching for versioned model names.
332        if let Some(&w) = table.get(model) {
333            return Some(w);
334        }
335        // Prefix fallback: match longest known key that is a prefix of `model`.
336        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    // ── private helpers ───────────────────────────────────────────────────────
344
345    fn build_context_table() -> HashMap<String, usize> {
346        let entries: &[(&str, usize)] = &[
347            // OpenAI GPT
348            ("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            // Anthropic Claude
362            ("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            // Google Gemini
374            ("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            // Meta LLaMA
381            ("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
391            ("mistral-7b", 32_768),
392            ("mistral-8x7b", 32_768),
393            ("mistral-large", 131_072),
394            ("mistral-small", 131_072),
395            // Cohere
396            ("command-r", 128_000),
397            ("command-r-plus", 128_000),
398            // DeepSeek
399            ("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            // Allow slight overshoot for single-sentence chunks larger than max.
447            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        // Message counting should be higher due to role overhead.
482        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}