mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(rust): count tool call arguments and tool_call_id like the python token counter
This commit is contained in:
parent
3cbb6ebc4a
commit
e06685c3e4
3 changed files with 63 additions and 4 deletions
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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"}}]}"#
|
||||
)]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue