From cc9cd492aaaf285f7c1286a763c4f24fb3e3a544 Mon Sep 17 00:00:00 2001 From: David Hagerty Date: Mon, 9 Feb 2026 10:51:14 -0500 Subject: [PATCH] feat(llm): add token usage tracking to Response type Add input_tokens and output_tokens fields to enable tracking of token consumption from LLM API calls. Each provider extracts token usage from their API responses: - Anthropic: input_tokens and output_tokens - OpenAI: prompt_tokens and completion_tokens - Ollama: prompt_eval_count and eval_count (None if unavailable) - Mock: supports configurable token counts Also add PartialEq to ResponseContent and ToolCall. Co-Authored-By: Claude Opus 4.6 --- src/llm/anthropic.rs | 12 ++++++ src/llm/mock.rs | 40 ++++++++++++------ src/llm/mod.rs | 6 ++- src/llm/ollama.rs | 14 +++++++ src/llm/openai.rs | 12 ++++++ tests/llm_token_tracking_test.rs | 70 ++++++++++++++++++++++++++++++++ 6 files changed, 139 insertions(+), 15 deletions(-) create mode 100644 tests/llm_token_tracking_test.rs diff --git a/src/llm/anthropic.rs b/src/llm/anthropic.rs index 881c897..f55c505 100644 --- a/src/llm/anthropic.rs +++ b/src/llm/anthropic.rs @@ -46,6 +46,16 @@ struct AnthropicTool { struct AnthropicResponse { content: Vec, stop_reason: Option, + #[serde(default)] + usage: UsageInfo, +} + +#[derive(Debug, Deserialize, Default)] +struct UsageInfo { + #[serde(default)] + input_tokens: usize, + #[serde(default)] + output_tokens: usize, } #[derive(Debug, Deserialize)] @@ -191,6 +201,8 @@ impl AnthropicClient { Ok(Response { content, stop_reason: anthropic_response.stop_reason, + input_tokens: Some(anthropic_response.usage.input_tokens), + output_tokens: Some(anthropic_response.usage.output_tokens), }) } } diff --git a/src/llm/mock.rs b/src/llm/mock.rs index b505ebc..5e76c26 100644 --- a/src/llm/mock.rs +++ b/src/llm/mock.rs @@ -6,8 +6,9 @@ use std::sync::{Arc, Mutex}; type RecordedCalls = Vec<(Vec, Vec)>; pub struct MockLlmClient { - responses: Arc>>, + responses: Arc)>>>, recorded_calls: Arc>, + token_counts: Arc>>, // (input_tokens, output_tokens) } impl MockLlmClient { @@ -15,27 +16,30 @@ impl MockLlmClient { Self { responses: Arc::new(Mutex::new(VecDeque::new())), recorded_calls: Arc::new(Mutex::new(Vec::new())), + token_counts: Arc::new(Mutex::new(None)), } } pub fn queue_text_response(&self, text: &str) { - let response = Response { - content: ResponseContent::Text(text.to_string()), - stop_reason: Some("end_turn".to_string()), - }; - self.responses.lock().unwrap().push_back(response); + self.responses.lock().unwrap().push_back(( + ResponseContent::Text(text.to_string()), + Some("end_turn".to_string()), + )); } pub fn queue_tool_call(&self, name: &str, params: serde_json::Value) { - let response = Response { - content: ResponseContent::ToolCalls(vec![ToolCall { + self.responses.lock().unwrap().push_back(( + ResponseContent::ToolCalls(vec![ToolCall { id: format!("call_{}", uuid::Uuid::new_v4()), name: name.to_string(), parameters: params, }]), - stop_reason: Some("tool_use".to_string()), - }; - self.responses.lock().unwrap().push_back(response); + Some("tool_use".to_string()), + )); + } + + pub fn set_token_counts(&self, input: usize, output: usize) { + *self.token_counts.lock().unwrap() = Some((input, output)); } pub fn get_recorded_calls(&self) -> Vec<(Vec, Vec)> { @@ -61,10 +65,20 @@ impl LlmClient for MockLlmClient { .unwrap() .push((messages, tools.to_vec())); - self.responses + let (content, stop_reason) = self + .responses .lock() .unwrap() .pop_front() - .ok_or_else(|| anyhow::anyhow!("No more mock responses queued")) + .ok_or_else(|| anyhow::anyhow!("No more mock responses queued"))?; + + let token_counts = *self.token_counts.lock().unwrap(); + + Ok(Response { + content, + stop_reason, + input_tokens: token_counts.map(|(i, _)| i), + output_tokens: token_counts.map(|(_, o)| o), + }) } } diff --git a/src/llm/mod.rs b/src/llm/mod.rs index a3fffa7..07ad1d1 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -64,7 +64,7 @@ pub struct ToolDefinition { } /// A tool call requested by the LLM -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] pub struct ToolCall { pub id: String, pub name: String, @@ -79,7 +79,7 @@ pub struct ToolResult { } /// Content of a response from the LLM -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] #[serde(untagged)] pub enum ResponseContent { Text(String), @@ -91,6 +91,8 @@ pub enum ResponseContent { pub struct Response { pub content: ResponseContent, pub stop_reason: Option, + pub input_tokens: Option, + pub output_tokens: Option, } /// Trait for LLM client implementations diff --git a/src/llm/ollama.rs b/src/llm/ollama.rs index 4a9ddd2..e6350f0 100644 --- a/src/llm/ollama.rs +++ b/src/llm/ollama.rs @@ -50,6 +50,10 @@ struct OllamaResponse { #[allow(dead_code)] done: bool, done_reason: Option, + #[serde(default)] + prompt_eval_count: usize, + #[serde(default)] + eval_count: usize, } #[derive(Debug, Deserialize)] @@ -181,6 +185,16 @@ impl OllamaClient { Ok(Response { content, stop_reason: ollama_response.done_reason, + input_tokens: if ollama_response.prompt_eval_count > 0 { + Some(ollama_response.prompt_eval_count) + } else { + None + }, + output_tokens: if ollama_response.eval_count > 0 { + Some(ollama_response.eval_count) + } else { + None + }, }) } } diff --git a/src/llm/openai.rs b/src/llm/openai.rs index d6282ea..3f699bc 100644 --- a/src/llm/openai.rs +++ b/src/llm/openai.rs @@ -51,6 +51,16 @@ struct OpenAiFunction { #[derive(Debug, Deserialize)] struct OpenAiResponse { choices: Vec, + #[serde(default)] + usage: OpenAiUsage, +} + +#[derive(Debug, Deserialize, Default)] +struct OpenAiUsage { + #[serde(default)] + prompt_tokens: usize, + #[serde(default)] + completion_tokens: usize, } #[derive(Debug, Deserialize)] @@ -191,6 +201,8 @@ impl OpenAiClient { Ok(Response { content, stop_reason: choice.finish_reason.clone(), + input_tokens: Some(openai_response.usage.prompt_tokens), + output_tokens: Some(openai_response.usage.completion_tokens), }) } } diff --git a/tests/llm_token_tracking_test.rs b/tests/llm_token_tracking_test.rs new file mode 100644 index 0000000..90629f6 --- /dev/null +++ b/tests/llm_token_tracking_test.rs @@ -0,0 +1,70 @@ +use rustagent::llm::mock::MockLlmClient; +use rustagent::llm::{LlmClient, Message, ResponseContent}; + +#[tokio::test] +async fn test_response_has_token_fields() { + // Verify Response struct contains input_tokens and output_tokens fields + let client = MockLlmClient::new(); + client.queue_text_response("Hello, world!"); + + let messages = vec![Message::user("Hi")]; + let response = client.chat(messages, &[]).await.unwrap(); + + // Fields should exist, though they may be None for mock responses + assert_eq!( + response.content, + ResponseContent::Text("Hello, world!".to_string()) + ); + // These fields should exist on Response + let _ = response.input_tokens; + let _ = response.output_tokens; +} + +#[tokio::test] +async fn test_mock_can_set_token_counts() { + // Verify mock client supports setting token counts + let client = MockLlmClient::new(); + client.queue_text_response("Hello, world!"); + client.set_token_counts(100, 50); + + let messages = vec![Message::user("Hi")]; + let response = client.chat(messages, &[]).await.unwrap(); + + assert_eq!(response.input_tokens, Some(100)); + assert_eq!(response.output_tokens, Some(50)); +} + +#[tokio::test] +async fn test_mock_token_counts_default_none() { + // Verify mock responses have None by default for token counts + let client = MockLlmClient::new(); + client.queue_text_response("Hello"); + + let messages = vec![Message::user("Hi")]; + let response = client.chat(messages, &[]).await.unwrap(); + + // Default should be None (no token info set) + assert_eq!(response.input_tokens, None); + assert_eq!(response.output_tokens, None); +} + +#[tokio::test] +async fn test_token_counts_with_tool_calls() { + // Verify token counts work with tool call responses too + let client = MockLlmClient::new(); + client.queue_tool_call("read_file", serde_json::json!({"path": "test.txt"})); + client.set_token_counts(75, 25); + + let messages = vec![Message::user("Read the file")]; + let response = client.chat(messages, &[]).await.unwrap(); + + assert_eq!(response.input_tokens, Some(75)); + assert_eq!(response.output_tokens, Some(25)); + match response.content { + ResponseContent::ToolCalls(calls) => { + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "read_file"); + } + _ => panic!("Expected tool call response"), + } +} -- 2.51.2