1use crate::message::{ContentPart, Message};
8use crate::tool_types::ToolCall;
9
10const CHARS_PER_TOKEN: usize = 4;
11
12#[must_use]
14pub fn estimate_tokens(text: &str) -> usize {
15 text.len().div_ceil(CHARS_PER_TOKEN)
16}
17
18#[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#[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#[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#[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}