From 25a6047a6127a4aae2a6a9df6d77d0d692dc52c9 Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Wed, 3 Jun 2026 10:12:53 -0400 Subject: [PATCH] refactor: move provider stream helpers Signed-off-by: Jack Amadeo --- crates/goose-providers/src/provider.rs | 148 ++++++++++++++++++++++++- crates/goose/src/providers/base.rs | 61 +--------- 2 files changed, 149 insertions(+), 60 deletions(-) diff --git a/crates/goose-providers/src/provider.rs b/crates/goose-providers/src/provider.rs index 7d3d08dfac26..43103da5d1bd 100644 --- a/crates/goose-providers/src/provider.rs +++ b/crates/goose-providers/src/provider.rs @@ -2,8 +2,8 @@ use crate::errors::ProviderError; use crate::models; use crate::retry::RetryConfig; use async_trait::async_trait; -use futures::Stream; -use goose_types::{Message, ModelConfig, ModelInfo, ProviderUsage}; +use futures::{Stream, StreamExt}; +use goose_types::{Message, MessageContent, ModelConfig, ModelInfo, ProviderUsage, Usage}; use rmcp::model::Tool; use std::pin::Pin; @@ -11,6 +11,59 @@ pub type MessageStream = Pin< Box, Option), ProviderError>> + Send>, >; +pub fn stream_from_single_message(message: Message, usage: ProviderUsage) -> MessageStream { + let stream = futures::stream::once(async move { Ok((Some(message), Some(usage))) }); + Box::pin(stream) +} + +pub async fn collect_stream( + mut stream: MessageStream, +) -> Result<(Message, ProviderUsage), ProviderError> { + let mut final_message: Option = None; + let mut final_usage: Option = None; + + while let Some(result) = stream.next().await { + let (msg_opt, usage_opt) = result?; + + if let Some(msg) = msg_opt { + final_message = Some(match final_message { + Some(mut prev) => { + for new_content in msg.content { + match (&mut prev.content.last_mut(), &new_content) { + ( + Some(MessageContent::Text(last_text)), + MessageContent::Text(new_text), + ) => { + last_text.text.push_str(&new_text.text); + } + _ => { + prev.content.push(new_content); + } + } + } + prev + } + None => msg, + }); + } + + if let Some(usage) = usage_opt { + final_usage = Some(usage); + } + } + + match final_message { + Some(msg) => { + let usage = final_usage + .unwrap_or_else(|| ProviderUsage::new("unknown".to_string(), Usage::default())); + Ok((msg, usage)) + } + None => Err(ProviderError::ExecutionError( + "Stream yielded no message".to_string(), + )), + } +} + #[async_trait] pub trait Provider: Send + Sync { fn get_name(&self) -> &str; @@ -113,3 +166,94 @@ pub trait Provider: Send + Sync { )) } } + +#[cfg(test)] +mod tests { + use super::*; + use futures::Stream; + use goose_types::{MessageContent, Usage}; + use rmcp::model::{CallToolRequestParams, Role}; + use test_case::test_case; + + fn content_from_str(s: String) -> MessageContent { + if let Some(img_data) = s.strip_prefix("*img:") { + MessageContent::image(format!("http://example.com/{img_data}"), "image/png") + } else if let Some(tool_name) = s.strip_prefix("*tool:") { + let tool_call = Ok(CallToolRequestParams::new(tool_name.to_string()) + .with_arguments(serde_json::Map::new())); + MessageContent::tool_request(format!("tool_{tool_name}"), tool_call) + } else { + MessageContent::text(s) + } + } + + fn create_test_stream( + items: Vec, + ) -> impl Stream, Option), ProviderError>> { + use futures::stream; + stream::iter(items.into_iter().map(|item| { + let content = content_from_str(item); + let message = Message::new(Role::Assistant, 0, vec![content]); + Ok((Some(message), None)) + })) + } + + fn content_to_strings(msg: &Message) -> Vec { + msg.content + .iter() + .map(|c| match c { + MessageContent::Text(t) => t.text.clone(), + MessageContent::Image(_) => "*img".to_string(), + MessageContent::ToolRequest(tr) => match &tr.tool_call { + Ok(call) => format!("*tool:{}", call.name), + Err(_) => "*tool:error".to_string(), + }, + _ => "*other".to_string(), + }) + .collect() + } + + #[test_case( + vec!["Hello", " ", "world"], + vec!["Hello world"] + ; "consecutive text coalesces" + )] + #[test_case( + vec!["Hello", "*img:pic1", "world"], + vec!["Hello", "*img", "world"] + ; "non-text breaks coalescing" + )] + #[test_case( + vec!["A", "B", "*img:pic1", "C", "D", "*tool:read", "E", "F"], + vec!["AB", "*img", "CD", "*tool:read", "EF"] + ; "multiple text groups" + )] + #[tokio::test] + async fn collect_stream_coalesces_text_content(input_items: Vec<&str>, expected: Vec<&str>) { + let items: Vec = input_items.into_iter().map(|s| s.to_string()).collect(); + let stream = create_test_stream(items); + let (msg, _) = collect_stream(Box::pin(stream)).await.unwrap(); + assert_eq!(content_to_strings(&msg), expected); + } + + #[tokio::test] + async fn collect_stream_defaults_usage() { + let stream = create_test_stream(vec!["Hello".to_string()]); + let (msg, usage) = collect_stream(Box::pin(stream)).await.unwrap(); + assert_eq!(content_to_strings(&msg), vec!["Hello"]); + assert_eq!(usage.model, "unknown"); + } + + #[tokio::test] + async fn stream_from_single_message_round_trips_message_and_usage() { + let message = Message::assistant().with_text("done"); + let usage = ProviderUsage::new("test-model".to_string(), Usage::default()); + + let (message, usage) = collect_stream(stream_from_single_message(message, usage)) + .await + .unwrap(); + + assert_eq!(message.as_concat_text(), "done"); + assert_eq!(usage.model, "test-model"); + } +} diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index 05f7c9e34922..1784b997b070 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -9,7 +9,7 @@ use super::inventory::{default_inventory_identity, InventoryIdentityInput}; use super::retry::RetryConfig; use crate::config::base::ConfigValue; use crate::config::{Config, ExtensionConfig, GooseMode}; -use crate::conversation::message::{Message, MessageContent}; +use crate::conversation::message::Message; use crate::conversation::Conversation; use crate::permission::PermissionConfirmation; use crate::utils::safe_truncate; @@ -22,6 +22,7 @@ use std::path::PathBuf; use std::sync::LazyLock; use std::sync::Mutex; +pub use goose_providers::provider::{collect_stream, stream_from_single_message}; pub use goose_providers::text::{split_think_blocks, FilterOut, ThinkFilter}; fn strip_xml_tags(text: &str) -> String { @@ -609,66 +610,10 @@ impl goose_providers::provider::Provider for GooseProviderAdapter<'_> { } } -pub fn stream_from_single_message(message: Message, usage: ProviderUsage) -> MessageStream { - let stream = futures::stream::once(async move { Ok((Some(message), Some(usage))) }); - Box::pin(stream) -} - -/// Collect all chunks from a MessageStream into a single Message and ProviderUsage -pub async fn collect_stream( - mut stream: MessageStream, -) -> Result<(Message, ProviderUsage), ProviderError> { - use futures::StreamExt; - - let mut final_message: Option = None; - let mut final_usage: Option = None; - - while let Some(result) = stream.next().await { - let (msg_opt, usage_opt) = result?; - - if let Some(msg) = msg_opt { - final_message = Some(match final_message { - Some(mut prev) => { - for new_content in msg.content { - match (&mut prev.content.last_mut(), &new_content) { - // Coalesce consecutive text blocks - ( - Some(MessageContent::Text(last_text)), - MessageContent::Text(new_text), - ) => { - last_text.text.push_str(&new_text.text); - } - _ => { - prev.content.push(new_content); - } - } - } - prev - } - None => msg, - }); - } - - if let Some(usage) = usage_opt { - final_usage = Some(usage); - } - } - - match final_message { - Some(msg) => { - let usage = final_usage - .unwrap_or_else(|| ProviderUsage::new("unknown".to_string(), Usage::default())); - Ok((msg, usage)) - } - None => Err(ProviderError::ExecutionError( - "Stream yielded no message".to_string(), - )), - } -} - #[cfg(test)] mod tests { use super::*; + use goose_types::MessageContent; use std::collections::HashMap; use test_case::test_case;