Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
148 changes: 146 additions & 2 deletions crates/goose-providers/src/provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,68 @@ 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;

pub type MessageStream = Pin<
Box<dyn Stream<Item = Result<(Option<Message>, Option<ProviderUsage>), 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<Message> = None;
let mut final_usage: Option<ProviderUsage> = 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;
Expand Down Expand Up @@ -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<String>,
) -> impl Stream<Item = Result<(Option<Message>, Option<ProviderUsage>), 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<String> {
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<String> = 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");
}
}
61 changes: 3 additions & 58 deletions crates/goose/src/providers/base.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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 {
Expand Down Expand Up @@ -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<Message> = None;
let mut final_usage: Option<ProviderUsage> = 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;

Expand Down