fix(rust): count tool call arguments and tool_call_id like the python token counter

This commit is contained in:
James 2026-09-23 14:16:39 -07:00
parent 3cbb6ebc4a
commit e06685c3e4
3 changed files with 63 additions and 4 deletions

View file

@ -124,7 +124,31 @@ impl TokenCounter {
.sum::<Result<usize, _>>()?,
None => 0,
};
Ok(TOKENS_PER_MESSAGE + role_tokens + name_tokens + content_tokens)
// Python counts only each tool call's `arguments` string
// (`_count_function_call_tokens`); names ride with the tool
// definitions and `tool_choice`.
let tool_call_tokens = match &message.tool_calls {
Some(calls) => calls
.iter()
.map(|call| self.count_text(call.function.arguments.as_deref().unwrap_or("")))
.sum::<Result<usize, _>>()?,
None => 0,
};
let tool_call_id_tokens = match &message.tool_call_id {
Some(id) => self.count_text(id)?,
None => 0,
};
let legacy_call_tokens = match &message.function_call {
Some(call) => self.count_text(call.arguments.as_deref().unwrap_or(""))?,
None => 0,
};
Ok(TOKENS_PER_MESSAGE
+ role_tokens
+ name_tokens
+ content_tokens
+ tool_call_tokens
+ tool_call_id_tokens
+ legacy_call_tokens)
}
fn count_content_item(&self, item: &ContentItem) -> Result<usize, Error> {

View file

@ -137,6 +137,28 @@ impl<'de> Visitor<'de> for TextValueVisitor {
}
}
/// The parts of an assistant tool call Python counts: only the `arguments`
/// string contributes (`_count_function_call_tokens` in
/// `litellm_core_utils/token_counter.py`). Absent arguments count as the empty
/// string; a non-string `arguments` declines so Python handles the fallback
/// instead of miscounting.
#[derive(Clone, Debug, Deserialize, PartialEq)]
pub(crate) struct ToolCall {
pub(crate) function: ToolCallFunction,
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
pub(crate) struct ToolCallFunction {
#[serde(default)]
pub(crate) arguments: Option<String>,
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
pub(crate) struct LegacyFunctionCall {
#[serde(default)]
pub(crate) arguments: Option<String>,
}
/// Python counts every string-valued key of a message, so any key beyond these
/// makes the shape unsupported rather than silently uncounted.
#[derive(Clone, Debug, Deserialize, PartialEq)]
@ -145,6 +167,9 @@ pub(crate) struct Message {
pub(crate) role: Option<String>,
pub(crate) name: Option<String>,
pub(crate) content: Option<MessageContent>,
pub(crate) tool_call_id: Option<String>,
pub(crate) tool_calls: Option<Vec<ToolCall>>,
pub(crate) function_call: Option<LegacyFunctionCall>,
}
#[derive(Clone, Debug, Deserialize, PartialEq)]

View file

@ -70,6 +70,16 @@ mod json {
const RERANK: &str = r#"{"model":"claude-sonnet-4-5","query":"best harbour",
"documents":["doc one",{"text":"doc two","title":"T","n":3,"ok":true,"none":null,"tags":["a","b"]}]}"#;
const TOOL_CALLS_AND_TOOL_RESULT: &str = r#"{"model":"claude-sonnet-4-5","messages":[
{"role":"assistant","tool_calls":[{"id":"1","type":"function","function":{"name":"f","arguments":"{\"city\": \"Tokyo\"}"}}]},
{"role":"tool","tool_call_id":"1","content":"Sunny, 22C"}]}"#;
const LEGACY_FUNCTION_CALL: &str = r#"{"model":"claude-sonnet-4-5","messages":[
{"role":"assistant","function_call":{"name":"f","arguments":"{\"city\": \"Tokyo\"}"}}]}"#;
const ASSISTANT_CONTENT_WITH_TOOL_CALL: &str = r#"{"model":"claude-sonnet-4-5","messages":[
{"role":"assistant","content":"calling","tool_calls":[{"id":"1","type":"function","function":{"name":"get_weather","arguments":"{\"city\": \"Paris, France\"}"}}]}]}"#;
fn assert_count_request_matches_python_token_counter(
load: JsonLoader,
body: &str,
@ -154,6 +164,9 @@ mod json {
#[case::responses_input_items(super::RESPONSES_INPUT, 62)]
#[case::embeddings_token_ids(super::EMBEDDINGS_TOKEN_IDS, 5)]
#[case::rerank_query_and_documents(super::RERANK, 41)]
#[case::tool_calls_and_tool_result(super::TOOL_CALLS_AND_TOOL_RESULT, 24)]
#[case::legacy_function_call(super::LEGACY_FUNCTION_CALL, 14)]
#[case::assistant_content_with_tool_call(super::ASSISTANT_CONTENT_WITH_TOOL_CALL, 16)]
fn count_request_matches_python_token_counter(
#[case] body: &str,
#[case] expected: usize,
@ -235,9 +248,6 @@ mod json {
#[rstest]
#[case::not_json(b"not json" as &[u8])]
#[case::messages_not_a_list(br#"{"model":"m","messages":"hi"}"#)]
#[case::message_with_tool_calls(
br#"{"model":"m","messages":[{"role":"assistant","tool_calls":[{"id":"1","type":"function","function":{"name":"f","arguments":"{}"}}]}]}"#
)]
#[case::dict_content(
br#"{"model":"m","messages":[{"role":"user","content":{"type":"text","text":"x"}}]}"#
)]