diff --git a/rust/agent/src/lib.rs b/rust/agent/src/lib.rs index d1e503dafcc..673c5191214 100644 --- a/rust/agent/src/lib.rs +++ b/rust/agent/src/lib.rs @@ -24,7 +24,7 @@ pub use inference::{ UnknownAnthropicModel, }; pub use provider::ProviderFormat; -pub use tool::{DynTool, Tool, ToolCallMetadata, ToolSet}; +pub use tool::{DynTool, SubagentUsage, Tool, ToolCallMetadata, ToolSet}; pub use tools::weather::{GetWeatherTool, TemperatureUnit}; pub use trajectory::{ Action, ActionBuilder, ActionItem, Call, Entry, Observation, ObservationBuilder, diff --git a/rust/agent/src/tool.rs b/rust/agent/src/tool.rs index a00af898c5f..527278dfa37 100644 --- a/rust/agent/src/tool.rs +++ b/rust/agent/src/tool.rs @@ -28,13 +28,17 @@ use crate::provider::ProviderFormat; pub enum ToolCallMetadata { /// Token usage reported by the deep-research subagent behind /// Foundation's `subagent_search` tool. - SubagentUsage { - model: String, - input_tokens: u64, - output_tokens: u64, - cache_read_tokens: u64, - cache_write_tokens: u64, - }, + SubagentUsage { usages: Vec }, +} + +/// Token usage from one model used by a deep-research subagent. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct SubagentUsage { + pub model: String, + pub input_tokens: u64, + pub output_tokens: u64, + pub cache_read_tokens: u64, + pub cache_write_tokens: u64, } /// THE trait you implement to define a tool. diff --git a/rust/foundation-api/src/agent_tools/subagent_search_tool.rs b/rust/foundation-api/src/agent_tools/subagent_search_tool.rs index 923d4ccfdca..3b00463e31c 100644 --- a/rust/foundation-api/src/agent_tools/subagent_search_tool.rs +++ b/rust/foundation-api/src/agent_tools/subagent_search_tool.rs @@ -10,7 +10,7 @@ use async_trait::async_trait; use schemars::JsonSchema; use serde::Deserialize; -use chroma_agent::{AgentError, Tool, ToolCallMetadata}; +use chroma_agent::{AgentError, SubagentUsage, Tool, ToolCallMetadata}; use crate::routes::subagent_search::{subagent_search_text, SubagentSearchCreds}; @@ -77,12 +77,17 @@ impl Tool for SubagentSearchTool { .await .map_err(|err| AgentError::Tool(err.to_string()))?; - let metadata = usage.map(|usage| ToolCallMetadata::SubagentUsage { - model: usage.model, - input_tokens: usage.input_tokens, - output_tokens: usage.output_tokens, - cache_read_tokens: usage.cache_read_tokens, - cache_write_tokens: usage.cache_write_tokens, + let metadata = (!usage.is_empty()).then(|| ToolCallMetadata::SubagentUsage { + usages: usage + .into_iter() + .map(|usage| SubagentUsage { + model: usage.model, + input_tokens: usage.input_tokens, + output_tokens: usage.output_tokens, + cache_read_tokens: usage.cache_read_tokens, + cache_write_tokens: usage.cache_write_tokens, + }) + .collect(), }); Ok((text, metadata)) diff --git a/rust/foundation-api/src/routes/agent/mod.rs b/rust/foundation-api/src/routes/agent/mod.rs index e5330d71930..82a1a70ea86 100644 --- a/rust/foundation-api/src/routes/agent/mod.rs +++ b/rust/foundation-api/src/routes/agent/mod.rs @@ -342,28 +342,27 @@ fn extract_subagent_usages(observation: &Observation) -> Vec { .iter() .filter_map(|item| { let ObservationItem::ToolResult { - metadata: - Some(chroma_agent::ToolCallMetadata::SubagentUsage { - model, - input_tokens, - output_tokens, - cache_read_tokens, - cache_write_tokens, - }), + metadata: Some(chroma_agent::ToolCallMetadata::SubagentUsage { usages }), .. } = item else { return None; }; - Some(InferenceUsage { - model: model.clone(), - input_tokens: *input_tokens, - output_tokens: *output_tokens, - cache_read_tokens: *cache_read_tokens, - cache_write_tokens: *cache_write_tokens, - }) + Some( + usages + .iter() + .map(|usage| InferenceUsage { + model: usage.model.clone(), + input_tokens: usage.input_tokens, + output_tokens: usage.output_tokens, + cache_read_tokens: usage.cache_read_tokens, + cache_write_tokens: usage.cache_write_tokens, + }) + .collect::>(), + ) }) + .flatten() .collect() } diff --git a/rust/foundation-api/src/routes/subagent_search/events.rs b/rust/foundation-api/src/routes/subagent_search/events.rs index f3daf3d913c..376234e4574 100644 --- a/rust/foundation-api/src/routes/subagent_search/events.rs +++ b/rust/foundation-api/src/routes/subagent_search/events.rs @@ -21,7 +21,7 @@ use std::sync::LazyLock; pub(crate) enum AgentEvent { Action(ActionData), Observation(ObservationData), - Usage(UsageData), + Usage(Vec), Done, Error(ErrorData), Unknown, @@ -41,7 +41,16 @@ impl AgentEvent { Some("observation") => { from_data(data()).map_or(AgentEvent::Unknown, AgentEvent::Observation) } - Some("usage") => from_data(data()).map_or(AgentEvent::Unknown, AgentEvent::Usage), + Some("usage") => match from_data::(data()) { + Some(usage) => AgentEvent::Usage(usage.into_records()), + None => { + tracing::warn!( + payload = %data(), + "dropping malformed search-agent usage event" + ); + AgentEvent::Unknown + } + }, Some("done") => AgentEvent::Done, Some("error") => from_data(data()).map_or(AgentEvent::Unknown, AgentEvent::Error), _ => AgentEvent::Unknown, @@ -102,6 +111,20 @@ pub(crate) struct UsageData { pub cache_write_tokens: u64, } +/// The usage envelope emitted by search-agent-research. +/// +/// Search agents report one record per model under `usage_records`. +#[derive(Debug, Clone, Deserialize)] +struct UsageEventData { + usage_records: Vec, +} + +impl UsageEventData { + fn into_records(self) -> Vec { + self.usage_records + } +} + impl ActionData { /// The `text` of this action's last `user_text` tool, if any — the agent /// "speaking" to the user. @@ -165,7 +188,7 @@ pub(crate) enum SubagentResultError { #[derive(Debug, Clone, PartialEq)] pub(crate) struct SubagentSearchResult { pub documents: Vec, - pub usage: Option, + pub usages: Vec, } /// Matches one `` diff --git a/rust/foundation-api/src/routes/subagent_search/mod.rs b/rust/foundation-api/src/routes/subagent_search/mod.rs index d432224cb36..f46d584f80c 100644 --- a/rust/foundation-api/src/routes/subagent_search/mod.rs +++ b/rust/foundation-api/src/routes/subagent_search/mod.rs @@ -359,12 +359,12 @@ pub(crate) async fn subagent_search_text( creds: SubagentSearchCreds, query: String, ui_origin: Option<&str>, -) -> Result<(String, Option), SubagentResultError> { +) -> Result<(String, Vec), SubagentResultError> { let tenant = creds.chroma_tenant.clone(); let result = collect_subagent_search_final(http, url, creds, query).await?; Ok(( format_ranked_documents(&result.documents, ui_origin, &tenant), - result.usage, + result.usages, )) } @@ -421,7 +421,7 @@ pub(crate) async fn collect_subagent_search_final( // Keep the last action's `user_text` — the agent's final answer. let mut final_answer: Option = None; - let mut usage: Option = None; + let mut usages = Vec::new(); let mut saw_done = false; while let Some(item) = stream.next().await { let raw = item.map_err(SubagentResultError::Stream)?; @@ -438,8 +438,8 @@ pub(crate) async fn collect_subagent_search_final( saw_done = true; break; } - AgentEvent::Usage(event_usage) => { - usage = Some(event_usage); + AgentEvent::Usage(event_usages) => { + usages.extend(event_usages); } AgentEvent::Observation(_) | AgentEvent::Unknown => {} } @@ -456,7 +456,7 @@ pub(crate) async fn collect_subagent_search_final( .as_deref() .map(parse_ranked_documents) .unwrap_or_default(), - usage, + usages, }) } diff --git a/rust/foundation-api/src/routes/subagent_search/tests/00_unit.rs b/rust/foundation-api/src/routes/subagent_search/tests/00_unit.rs index 1397c7a30ab..d5802702e32 100644 --- a/rust/foundation-api/src/routes/subagent_search/tests/00_unit.rs +++ b/rust/foundation-api/src/routes/subagent_search/tests/00_unit.rs @@ -40,15 +40,34 @@ fn agent_event_parses_each_kind() { )); assert!(matches!( AgentEvent::parse( - &json!({"type":"usage","data":{"model":"scout","input_tokens":123,"output_tokens":456}}).to_string() + &json!({ + "type": "usage", + "data": { + "usage_records": [ + {"model":"scout","input_tokens":123,"output_tokens":456}, + {"model":"max","input_tokens":7,"output_tokens":8,"cache_read_tokens":9} + ] + } + }) + .to_string() ), - AgentEvent::Usage(UsageData { - model, - input_tokens: 123, - output_tokens: 456, - cache_read_tokens: 0, - cache_write_tokens: 0, - }) if model == "scout" + AgentEvent::Usage(usages) + if usages == vec![ + UsageData { + model: "scout".to_string(), + input_tokens: 123, + output_tokens: 456, + cache_read_tokens: 0, + cache_write_tokens: 0, + }, + UsageData { + model: "max".to_string(), + input_tokens: 7, + output_tokens: 8, + cache_read_tokens: 9, + cache_write_tokens: 0, + }, + ] )); assert!(matches!( AgentEvent::parse(&json!({"type":"done","data":{}}).to_string()), @@ -66,6 +85,20 @@ fn agent_event_parses_each_kind() { assert!(matches!(AgentEvent::parse("not json"), AgentEvent::Unknown)); } +#[test] +fn malformed_usage_events_are_unknown() { + for data in [ + json!({}), + json!({"usage_record": []}), + json!({"model": "scout", "input_tokens": 123, "output_tokens": 456}), + ] { + assert!(matches!( + AgentEvent::parse(&json!({"type": "usage", "data": data}).to_string()), + AgentEvent::Unknown + )); + } +} + #[test] fn action_user_text_takes_last_and_detects_answer_only() { // A tool-call action with no user_text. diff --git a/rust/foundation-api/src/routes/subagent_search/tests/01_integration.rs b/rust/foundation-api/src/routes/subagent_search/tests/01_integration.rs index d205f22586d..db8ec4c8024 100644 --- a/rust/foundation-api/src/routes/subagent_search/tests/01_integration.rs +++ b/rust/foundation-api/src/routes/subagent_search/tests/01_integration.rs @@ -24,7 +24,7 @@ async fn streams_and_collects_final_from_mocked_sse() { "data: {\"type\":\"action\",\"data\":{\"tools\":[{\"name\":\"search\"}],\"params\":[{\"query\":\"rag\"}]}}\n\n", "data: {\"type\":\"observation\",\"data\":{\"sources\":[\"a\"]}}\n\n", "data: {\"type\":\"action\",\"data\":{\"tools\":[{\"name\":\"user_text\"}],\"params\":[{\"text\":\"Relevant to rag.\"}]}}\n\n", - "data: {\"type\":\"usage\",\"data\":{\"model\":\"scout\",\"input_tokens\":123,\"output_tokens\":456}}\n\n", + "data: {\"type\":\"usage\",\"data\":{\"usage_records\":[{\"model\":\"scout\",\"input_tokens\":123,\"output_tokens\":456},{\"model\":\"max\",\"input_tokens\":7,\"output_tokens\":8,\"cache_read_tokens\":9}]}}\n\n", "data: {\"type\":\"done\",\"data\":{}}\n\n", ); let mock = server @@ -69,12 +69,14 @@ async fn streams_and_collects_final_from_mocked_sse() { justification: "Relevant to rag.".to_string(), }] ); - let usage = result.usage.expect("usage should be present"); - assert_eq!(usage.model, "scout"); - assert_eq!(usage.input_tokens, 123); - assert_eq!(usage.output_tokens, 456); - assert_eq!(usage.cache_read_tokens, 0); - assert_eq!(usage.cache_write_tokens, 0); + assert_eq!(result.usages.len(), 2); + assert_eq!(result.usages[0].model, "scout"); + assert_eq!(result.usages[0].input_tokens, 123); + assert_eq!(result.usages[0].output_tokens, 456); + assert_eq!(result.usages[0].cache_read_tokens, 0); + assert_eq!(result.usages[0].cache_write_tokens, 0); + assert_eq!(result.usages[1].model, "max"); + assert_eq!(result.usages[1].cache_read_tokens, 9); assert_eq!(mock.calls(), 2); } @@ -180,7 +182,7 @@ async fn answer_with_no_documents_emits_empty_result_then_done() { .await .expect("empty results are Ok"); assert!(result.documents.is_empty()); - assert!(result.usage.is_none()); + assert!(result.usages.is_empty()); assert_eq!(mock.calls(), 2); }