Skip to main content

behest_core/
token.rs

1//! Token estimation using character-based heuristics.
2//!
3//! Uses the `chars / 4` rule-of-thumb to estimate token counts without
4//! requiring a tokenizer. The 20,000-token buffer in compaction absorbs
5//! estimation error, and the heuristic is consistent with industry practice.
6
7use crate::message::{ContentPart, Message};
8use crate::tool_types::ToolCall;
9
10const CHARS_PER_TOKEN: usize = 4;
11
12/// Estimates the number of tokens for a string.
13#[must_use]
14pub fn estimate_tokens(text: &str) -> usize {
15    text.len().div_ceil(CHARS_PER_TOKEN)
16}
17
18/// Estimates the token count for a content part.
19#[must_use]
20pub fn estimate_content_part_tokens(part: &ContentPart) -> usize {
21    match part {
22        ContentPart::Text { text } => estimate_tokens(text),
23        ContentPart::Json { value } => estimate_tokens(&value.to_string()),
24        ContentPart::ImageUrl { url, .. } => estimate_tokens(url),
25    }
26}
27
28/// Estimates the token count for a tool call.
29#[must_use]
30pub fn estimate_tool_call_tokens(call: &ToolCall) -> usize {
31    let name_tokens = estimate_tokens(&call.name);
32    let args_tokens = estimate_tokens(&call.arguments.to_string());
33    name_tokens + args_tokens + 20
34}
35
36/// Estimates the total token count for a provider [`Message`].
37#[must_use]
38pub fn estimate_message_tokens(message: &Message) -> usize {
39    match message {
40        Message::System { content } | Message::User { content } => {
41            content
42                .iter()
43                .map(estimate_content_part_tokens)
44                .sum::<usize>()
45                + 8
46        }
47        Message::Assistant {
48            content,
49            tool_calls,
50        } => {
51            let content_tokens: usize = content.iter().map(estimate_content_part_tokens).sum();
52            let tool_tokens: usize = tool_calls.iter().map(estimate_tool_call_tokens).sum();
53            content_tokens + tool_tokens + 8
54        }
55        Message::Tool {
56            tool_call_id,
57            name,
58            content,
59        } => {
60            let id_tokens = estimate_tokens(tool_call_id);
61            let name_tokens = estimate_tokens(name);
62            let content_tokens: usize = content.iter().map(estimate_content_part_tokens).sum();
63            id_tokens + name_tokens + content_tokens + 10
64        }
65    }
66}
67
68/// Estimates the total token count for a slice of provider [`Message`]s.
69#[must_use]
70pub fn estimate_messages_tokens(messages: &[Message]) -> usize {
71    messages.iter().map(estimate_message_tokens).sum()
72}
73
74#[cfg(test)]
75#[allow(clippy::unwrap_used)]
76mod tests {
77    use super::*;
78    use serde_json::json;
79
80    #[test]
81    fn estimate_should_round_up_for_fractional_tokens() {
82        assert_eq!(estimate_tokens(""), 0);
83        assert_eq!(estimate_tokens("a"), 1);
84        assert_eq!(estimate_tokens("ab"), 1);
85        assert_eq!(estimate_tokens("abcd"), 1);
86        assert_eq!(estimate_tokens("abcde"), 2);
87        assert_eq!(estimate_tokens("12345678"), 2);
88    }
89
90    #[test]
91    fn estimate_content_part_text() {
92        let part = ContentPart::text("Hello, world!");
93        assert_eq!(estimate_content_part_tokens(&part), 4);
94    }
95
96    #[test]
97    fn estimate_content_part_json() {
98        let part = ContentPart::json(json!({"key": "value"}));
99        assert_eq!(estimate_content_part_tokens(&part), 4);
100    }
101
102    #[test]
103    fn estimate_message_system() {
104        let msg = Message::system_text("You are helpful.");
105        assert_eq!(estimate_message_tokens(&msg), 12);
106    }
107
108    #[test]
109    fn estimate_message_user() {
110        let msg = Message::user_text("Hello");
111        assert_eq!(estimate_message_tokens(&msg), 10);
112    }
113
114    #[test]
115    fn estimate_messages_slice() {
116        let messages = vec![
117            Message::system_text("System"),
118            Message::user_text("User"),
119            Message::assistant_text("Assistant"),
120        ];
121        let total = estimate_messages_tokens(&messages);
122        assert!(total > 0);
123    }
124
125    #[test]
126    fn estimate_tool_call_includes_overhead() {
127        let call = ToolCall::new("call_1", "echo", json!({}));
128        let tokens = estimate_tool_call_tokens(&call);
129        assert_eq!(tokens, 22);
130    }
131}