From 1c820addebf4288c40bdd01f31ca940bdc05f4d4 Mon Sep 17 00:00:00 2001 From: Douwe Osinga Date: Tue, 16 Dec 2025 13:24:53 -0500 Subject: [PATCH] OpenRouter & Xai streaming (#5873) Co-authored-by: Douwe Osinga --- Cargo.lock | 11 +- crates/goose/src/providers/azure.rs | 9 +- crates/goose/src/providers/databricks.rs | 28 +--- crates/goose/src/providers/formats/openai.rs | 57 ++++--- crates/goose/src/providers/githubcopilot.rs | 9 +- crates/goose/src/providers/litellm.rs | 1 + crates/goose/src/providers/ollama.rs | 38 +---- crates/goose/src/providers/openai.rs | 70 ++++----- crates/goose/src/providers/openrouter.rs | 56 ++++++- crates/goose/src/providers/tetrate.rs | 58 +++---- crates/goose/src/providers/toolshim.rs | 1 + crates/goose/src/providers/utils.rs | 51 +++++-- crates/goose/src/providers/xai.rs | 52 ++++++- .../settings/app/ExternalBackendSection.tsx | 142 +++++++++--------- 14 files changed, 334 insertions(+), 249 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index bd4a6be5e1..23aea46f22 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4123,6 +4123,15 @@ dependencies = [ "either", ] +[[package]] +name = "itertools" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "413ee7dfc52ee1a4949ceeb7dbc8a33f2d6c088194d9f922fb8318faf1f01186" +dependencies = [ + "either", +] + [[package]] name = "itertools" version = "0.14.0" @@ -5695,7 +5704,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" dependencies = [ "anyhow", - "itertools 0.14.0", + "itertools 0.13.0", "proc-macro2", "quote", "syn 2.0.111", diff --git a/crates/goose/src/providers/azure.rs b/crates/goose/src/providers/azure.rs index d519a82c73..346ebeecf0 100644 --- a/crates/goose/src/providers/azure.rs +++ b/crates/goose/src/providers/azure.rs @@ -149,7 +149,14 @@ impl Provider for AzureProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; + let payload = create_request( + model_config, + system, + messages, + tools, + &ImageFormat::OpenAi, + false, + )?; let response = self .with_retry(|| async { let payload_clone = payload.clone(); diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index e4d3b3cf1a..6baa686d29 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -1,13 +1,8 @@ use anyhow::Result; -use async_stream::try_stream; use async_trait::async_trait; -use futures::TryStreamExt; use serde::{Deserialize, Serialize}; use serde_json::Value; -use std::io; use std::time::Duration; -use tokio::pin; -use tokio_util::io::StreamReader; use super::api_client::{ApiClient, AuthMethod, AuthProvider}; use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage}; @@ -17,21 +12,19 @@ use super::formats::databricks::{create_request, response_to_message}; use super::oauth; use super::retry::ProviderRetry; use super::utils::{ - get_model, handle_response_openai_compat, map_http_error_to_provider_error, ImageFormat, - RequestLog, + get_model, handle_response_openai_compat, map_http_error_to_provider_error, + stream_openai_compat, ImageFormat, RequestLog, }; use crate::config::ConfigError; use crate::conversation::message::Message; use crate::model::ModelConfig; -use crate::providers::formats::openai::{get_usage, response_to_streaming_message}; +use crate::providers::formats::openai::get_usage; use crate::providers::retry::{ RetryConfig, DEFAULT_BACKOFF_MULTIPLIER, DEFAULT_INITIAL_RETRY_INTERVAL_MS, DEFAULT_MAX_RETRIES, DEFAULT_MAX_RETRY_INTERVAL_MS, }; use rmcp::model::Tool; use serde_json::json; -use tokio_stream::StreamExt; -use tokio_util::codec::{FramedRead, LinesCodec}; const DEFAULT_CLIENT_ID: &str = "databricks-cli"; const DEFAULT_REDIRECT_URL: &str = "http://localhost"; @@ -347,20 +340,7 @@ impl Provider for DatabricksProvider { let _ = log.error(e); })?; - let stream = response.bytes_stream().map_err(io::Error::other); - - Ok(Box::pin(try_stream! { - let stream_reader = StreamReader::new(stream); - let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from); - - let message_stream = response_to_streaming_message(framed); - pin!(message_stream); - while let Some(message) = message_stream.next().await { - let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; - log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?; - yield (message, usage); - } - })) + stream_openai_compat(response, log) } fn supports_streaming(&self) -> bool { diff --git a/crates/goose/src/providers/formats/openai.rs b/crates/goose/src/providers/formats/openai.rs index 04fd962a1b..937d9658ff 100644 --- a/crates/goose/src/providers/formats/openai.rs +++ b/crates/goose/src/providers/formats/openai.rs @@ -613,6 +613,7 @@ pub fn create_request( messages: &[Message], tools: &[Tool], image_format: &ImageFormat, + for_streaming: bool, ) -> anyhow::Result { if model_config.model_name.starts_with("o1-mini") { return Err(anyhow!( @@ -652,13 +653,8 @@ pub fn create_request( }); let messages_spec = format_messages(messages, image_format); - let mut tools_spec = if !tools.is_empty() { - format_tools(tools)? - } else { - vec![] - }; + let mut tools_spec = format_tools(tools)?; - // Validate tool schemas validate_tool_schemas(&mut tools_spec); let mut messages_array = vec![system_message]; @@ -670,25 +666,17 @@ pub fn create_request( }); if let Some(effort) = reasoning_effort { - payload - .as_object_mut() - .unwrap() - .insert("reasoning_effort".to_string(), json!(effort)); + payload["reasoning_effort"] = json!(effort); } if !tools_spec.is_empty() { - payload - .as_object_mut() - .unwrap() - .insert("tools".to_string(), json!(tools_spec)); + payload["tools"] = json!(tools_spec); } + // o1, o3 models currently don't support temperature if !is_ox_model { if let Some(temp) = model_config.temperature { - payload - .as_object_mut() - .unwrap() - .insert("temperature".to_string(), json!(temp)); + payload["temperature"] = json!(temp); } } @@ -704,6 +692,12 @@ pub fn create_request( .unwrap() .insert(key.to_string(), json!(tokens)); } + + if for_streaming { + payload["stream"] = json!(true); + payload["stream_options"] = json!({"include_usage": true}); + } + Ok(payload) } @@ -1277,7 +1271,14 @@ mod tests { toolshim_model: None, fast_model: None, }; - let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; + let request = create_request( + &model_config, + "system", + &[], + &[], + &ImageFormat::OpenAi, + false, + )?; let obj = request.as_object().unwrap(); let expected = json!({ "model": "gpt-4o", @@ -1309,7 +1310,14 @@ mod tests { toolshim_model: None, fast_model: None, }; - let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; + let request = create_request( + &model_config, + "system", + &[], + &[], + &ImageFormat::OpenAi, + false, + )?; let obj = request.as_object().unwrap(); let expected = json!({ "model": "o1", @@ -1342,7 +1350,14 @@ mod tests { toolshim_model: None, fast_model: None, }; - let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; + let request = create_request( + &model_config, + "system", + &[], + &[], + &ImageFormat::OpenAi, + false, + )?; let obj = request.as_object().unwrap(); let expected = json!({ "model": "o3-mini", diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index d10dbf2aed..681afd4ef5 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -460,7 +460,14 @@ impl Provider for GithubCopilotProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; + let payload = create_request( + model_config, + system, + messages, + tools, + &ImageFormat::OpenAi, + false, + )?; let mut log = RequestLog::start(model_config, &payload)?; // Make request with retry diff --git a/crates/goose/src/providers/litellm.rs b/crates/goose/src/providers/litellm.rs index cc4bcd5046..308740e099 100644 --- a/crates/goose/src/providers/litellm.rs +++ b/crates/goose/src/providers/litellm.rs @@ -179,6 +179,7 @@ impl Provider for LiteLLMProvider { messages, tools, &ImageFormat::OpenAi, + false, )?; if self.supports_cache_control().await { diff --git a/crates/goose/src/providers/ollama.rs b/crates/goose/src/providers/ollama.rs index fc06ae4253..be1d927ff9 100644 --- a/crates/goose/src/providers/ollama.rs +++ b/crates/goose/src/providers/ollama.rs @@ -3,7 +3,8 @@ use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, Provider use super::errors::ProviderError; use super::retry::ProviderRetry; use super::utils::{ - get_model, handle_response_openai_compat, handle_status_openai_compat, RequestLog, + get_model, handle_response_openai_compat, handle_status_openai_compat, stream_openai_compat, + RequestLog, }; use crate::config::declarative_providers::DeclarativeProviderConfig; use crate::config::GooseMode; @@ -11,23 +12,14 @@ use crate::conversation::message::Message; use crate::conversation::Conversation; use crate::model::ModelConfig; -use crate::providers::formats::openai::{ - create_request, get_usage, response_to_message, response_to_streaming_message, -}; +use crate::providers::formats::openai::{create_request, get_usage, response_to_message}; use crate::utils::safe_truncate; use anyhow::Result; -use async_stream::try_stream; use async_trait::async_trait; -use futures::TryStreamExt; use regex::Regex; use rmcp::model::Tool; -use serde_json::{json, Value}; -use std::io; +use serde_json::Value; use std::time::Duration; -use tokio::pin; -use tokio_stream::StreamExt; -use tokio_util::codec::{FramedRead, LinesCodec}; -use tokio_util::io::StreamReader; use url::Url; pub const OLLAMA_HOST: &str = "localhost"; @@ -200,6 +192,7 @@ impl Provider for OllamaProvider { messages, filtered_tools, &super::utils::ImageFormat::OpenAi, + false, )?; let mut log = RequestLog::start(model_config, &payload)?; @@ -262,17 +255,14 @@ impl Provider for OllamaProvider { tools }; - let mut payload = create_request( + let payload = create_request( &self.model, system, messages, filtered_tools, &super::utils::ImageFormat::OpenAi, + true, )?; - payload["stream"] = json!(true); - payload["stream_options"] = json!({ - "include_usage": true, - }); let mut log = RequestLog::start(&self.model, &payload)?; let response = self @@ -287,19 +277,7 @@ impl Provider for OllamaProvider { .inspect_err(|e| { let _ = log.error(e); })?; - let stream = response.bytes_stream().map_err(io::Error::other); - - Ok(Box::pin(try_stream! { - let stream_reader = StreamReader::new(stream); - let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from); - let message_stream = response_to_streaming_message(framed); - pin!(message_stream); - while let Some(message) = message_stream.next().await { - let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; - log.write(&message, usage.as_ref().map(|f| &f.usage))?; - yield (message, usage); - } - })) + stream_openai_compat(response, log) } async fn fetch_supported_models(&self) -> Result>, ProviderError> { diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index 9af5579cd1..cf98dcbb53 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -1,33 +1,30 @@ -use anyhow::Result; -use async_stream::try_stream; -use async_trait::async_trait; -use futures::TryStreamExt; -use reqwest::StatusCode; -use serde_json::{json, Value}; -use std::collections::HashMap; -use std::io; -use tokio::pin; -use tokio_stream::StreamExt; -use tokio_util::codec::{FramedRead, LinesCodec}; -use tokio_util::io::StreamReader; - use super::api_client::{ApiClient, AuthMethod}; use super::base::{ConfigKey, ModelInfo, Provider, ProviderMetadata, ProviderUsage, Usage}; use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse}; use super::errors::ProviderError; -use super::formats::openai::{ - create_request, get_usage, response_to_message, response_to_streaming_message, -}; +use super::formats::openai::{create_request, get_usage, response_to_message}; use super::formats::openai_responses::{ create_responses_request, get_responses_usage, responses_api_to_message, responses_api_to_streaming_message, ResponsesApiResponse, }; use super::retry::ProviderRetry; use super::utils::{ - get_model, handle_response_openai_compat, handle_status_openai_compat, ImageFormat, + get_model, handle_response_openai_compat, handle_status_openai_compat, stream_openai_compat, + ImageFormat, }; use crate::config::declarative_providers::DeclarativeProviderConfig; use crate::conversation::message::Message; +use anyhow::Result; +use async_stream::try_stream; +use async_trait::async_trait; +use futures::{StreamExt, TryStreamExt}; +use reqwest::StatusCode; +use serde_json::Value; +use std::collections::HashMap; +use std::io; +use tokio::pin; +use tokio_util::codec::{FramedRead, LinesCodec}; +use tokio_util::io::StreamReader; use crate::model::ModelConfig; use crate::providers::base::MessageStream; @@ -286,8 +283,14 @@ impl Provider for OpenAiProvider { log.write(&json_response, Some(&usage))?; Ok((message, ProviderUsage::new(model, usage))) } else { - let payload = - create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; + let payload = create_request( + model_config, + system, + messages, + tools, + &ImageFormat::OpenAi, + false, + )?; let mut log = RequestLog::start(&self.model, &payload)?; let json_response = self @@ -404,12 +407,14 @@ impl Provider for OpenAiProvider { } })) } else { - let mut payload = - create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?; - payload["stream"] = serde_json::Value::Bool(true); - payload["stream_options"] = json!({ - "include_usage": true, - }); + let payload = create_request( + &self.model, + system, + messages, + tools, + &ImageFormat::OpenAi, + true, + )?; let mut log = RequestLog::start(&self.model, &payload)?; let response = self @@ -425,20 +430,7 @@ impl Provider for OpenAiProvider { let _ = log.error(e); })?; - let stream = response.bytes_stream().map_err(io::Error::other); - - Ok(Box::pin(try_stream! { - let stream_reader = StreamReader::new(stream); - let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from); - - let message_stream = response_to_streaming_message(framed); - pin!(message_stream); - while let Some(message) = message_stream.next().await { - let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; - log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?; - yield (message, usage); - } - })) + stream_openai_compat(response, log) } } } diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index 9869584d1b..23b8a7422d 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -3,12 +3,12 @@ use async_trait::async_trait; use serde_json::{json, Value}; use super::api_client::{ApiClient, AuthMethod}; -use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage}; +use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage}; use super::errors::ProviderError; use super::retry::ProviderRetry; use super::utils::{ - get_model, handle_response_google_compat, handle_response_openai_compat, is_google_model, - RequestLog, + get_model, handle_response_google_compat, handle_response_openai_compat, + handle_status_openai_compat, is_google_model, stream_openai_compat, RequestLog, }; use crate::conversation::message::Message; @@ -40,6 +40,7 @@ pub struct OpenRouterProvider { #[serde(skip)] api_client: ApiClient, model: ModelConfig, + supports_streaming: bool, #[serde(skip)] name: String, } @@ -62,6 +63,7 @@ impl OpenRouterProvider { Ok(Self { api_client, model, + supports_streaming: true, name: Self::metadata().name, }) } @@ -211,13 +213,13 @@ async fn create_request_based_on_model( messages, tools, &super::utils::ImageFormat::OpenAi, + false, )?; if provider.supports_cache_control().await { payload = update_request_for_anthropic(&payload); } - // Always add transforms: ["middle-out"] for OpenRouter to handle prompts > context size payload .as_object_mut() .unwrap() @@ -371,4 +373,50 @@ impl Provider for OpenRouterProvider { .model_name .starts_with(OPENROUTER_MODEL_PREFIX_ANTHROPIC) } + + fn supports_streaming(&self) -> bool { + self.supports_streaming + } + + async fn stream( + &self, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result { + let mut payload = create_request( + &self.model, + system, + messages, + tools, + &super::utils::ImageFormat::OpenAi, + true, + )?; + + if self.supports_cache_control().await { + payload = update_request_for_anthropic(&payload); + } + + payload + .as_object_mut() + .unwrap() + .insert("transforms".to_string(), json!(["middle-out"])); + + let mut log = RequestLog::start(&self.model, &payload)?; + + let response = self + .with_retry(|| async { + let resp = self + .api_client + .response_post("api/v1/chat/completions", &payload) + .await?; + handle_status_openai_compat(resp).await + }) + .await + .inspect_err(|e| { + let _ = log.error(e); + })?; + + stream_openai_compat(response, log) + } } diff --git a/crates/goose/src/providers/tetrate.rs b/crates/goose/src/providers/tetrate.rs index 9547628d46..ab3e7af0ca 100644 --- a/crates/goose/src/providers/tetrate.rs +++ b/crates/goose/src/providers/tetrate.rs @@ -1,25 +1,16 @@ -use anyhow::Result; -use async_stream::try_stream; -use async_trait::async_trait; -use futures::TryStreamExt; -use serde_json::{json, Value}; -use std::io; -use tokio::pin; -use tokio_stream::StreamExt; -use tokio_util::codec::{FramedRead, LinesCodec}; -use tokio_util::io::StreamReader; - use super::api_client::{ApiClient, AuthMethod}; use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage}; use super::errors::ProviderError; -use super::formats::openai::response_to_streaming_message; use super::retry::ProviderRetry; use super::utils::{ get_model, handle_response_google_compat, handle_response_openai_compat, - handle_status_openai_compat, is_google_model, RequestLog, + handle_status_openai_compat, is_google_model, stream_openai_compat, RequestLog, }; use crate::config::signup_tetrate::TETRATE_DEFAULT_MODEL; use crate::conversation::message::Message; +use anyhow::Result; +use async_trait::async_trait; +use serde_json::Value; use crate::model::ModelConfig; use crate::providers::formats::openai::{create_request, get_usage, response_to_message}; @@ -178,6 +169,7 @@ impl Provider for TetrateProvider { messages, tools, &super::utils::ImageFormat::OpenAi, + false, )?; let mut log = RequestLog::start(model_config, &payload)?; @@ -206,41 +198,31 @@ impl Provider for TetrateProvider { messages: &[Message], tools: &[Tool], ) -> Result { - let mut payload = create_request( + let payload = create_request( &self.model, system, messages, tools, &super::utils::ImageFormat::OpenAi, + true, )?; - payload["stream"] = json!(true); - payload["stream_options"] = json!({ - "include_usage": true, - }); - - let resp = self - .api_client - .response_post("v1/chat/completions", &payload) - .await?; - - let response = handle_status_openai_compat(resp).await?; - - let stream = response.bytes_stream().map_err(io::Error::other); let mut log = RequestLog::start(&self.model, &payload)?; - Ok(Box::pin(try_stream! { - let stream_reader = StreamReader::new(stream); - let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from); + let response = self + .with_retry(|| async { + let resp = self + .api_client + .response_post("v1/chat/completions", &payload) + .await?; + handle_status_openai_compat(resp).await + }) + .await + .inspect_err(|e| { + let _ = log.error(e); + })?; - let message_stream = response_to_streaming_message(framed); - pin!(message_stream); - while let Some(message) = message_stream.next().await { - let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; - log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?; - yield (message, usage); - } - })) + stream_openai_compat(response, log) } /// Fetch supported models from Tetrate Agent Router Service API (only models with tool support) diff --git a/crates/goose/src/providers/toolshim.rs b/crates/goose/src/providers/toolshim.rs index b4dc3cb11b..325f6da5a2 100644 --- a/crates/goose/src/providers/toolshim.rs +++ b/crates/goose/src/providers/toolshim.rs @@ -162,6 +162,7 @@ impl OllamaInterpreter { &messages, &[], // No tools &super::utils::ImageFormat::OpenAi, + false, )?; payload["stream"] = json!(false); // needed for the /api/chat endpoint to work diff --git a/crates/goose/src/providers/utils.rs b/crates/goose/src/providers/utils.rs index d059d6695e..6a2d25f87d 100644 --- a/crates/goose/src/providers/utils.rs +++ b/crates/goose/src/providers/utils.rs @@ -1,10 +1,13 @@ -use super::base::Usage; +use super::base::{MessageStream, Usage}; use super::errors::GoogleErrorCode; use crate::config::paths::Paths; use crate::model::ModelConfig; use crate::providers::errors::ProviderError; +use crate::providers::formats::openai::response_to_streaming_message; use anyhow::{anyhow, Result}; +use async_stream::try_stream; use base64::Engine; +use futures::TryStreamExt; use regex::Regex; use reqwest::{Response, StatusCode}; use rmcp::model::{AnnotateAble, ImageContent, RawImageContent}; @@ -12,9 +15,14 @@ use serde::{Deserialize, Serialize}; use serde_json::{json, Map, Value}; use std::fmt::Display; use std::fs::File; +use std::io; use std::io::{BufWriter, Read, Write}; use std::path::{Path, PathBuf}; use std::time::Duration; +use tokio::pin; +use tokio_stream::StreamExt; +use tokio_util::codec::{FramedRead, LinesCodec}; +use tokio_util::io::StreamReader; use uuid::Uuid; #[derive(Debug, Copy, Clone, Serialize, Deserialize)] @@ -178,19 +186,36 @@ pub async fn handle_response_openai_compat(response: Response) -> Result Result { + let stream = response.bytes_stream().map_err(io::Error::other); + + Ok(Box::pin(try_stream! { + let stream_reader = StreamReader::new(stream); + let framed = FramedRead::new(stream_reader, LinesCodec::new()) + .map_err(anyhow::Error::from); + + let message_stream = response_to_streaming_message(framed); + pin!(message_stream); + while let Some(message) = message_stream.next().await { + let (message, usage) = message.map_err(|e| + ProviderError::RequestFailed(format!("Stream decode error: {}", e)) + )?; + log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?; + yield (message, usage); + } + })) +} + pub fn is_google_model(payload: &Value) -> bool { - if let Some(model) = payload.get("model").and_then(|m| m.as_str()) { - // Check if the model name contains "google" - return model.to_lowercase().contains("google"); - } - false + payload + .get("model") + .and_then(|m| m.as_str()) + .unwrap_or("") + .to_lowercase() + .contains("google") } /// Extracts `StatusCode` from response status or payload error code. diff --git a/crates/goose/src/providers/xai.rs b/crates/goose/src/providers/xai.rs index 0078d9894e..b151aaf184 100644 --- a/crates/goose/src/providers/xai.rs +++ b/crates/goose/src/providers/xai.rs @@ -1,17 +1,20 @@ use super::api_client::{ApiClient, AuthMethod}; use super::errors::ProviderError; use super::retry::ProviderRetry; -use super::utils::{get_model, handle_response_openai_compat, RequestLog}; +use super::utils::{ + get_model, handle_response_openai_compat, handle_status_openai_compat, stream_openai_compat, + RequestLog, +}; use crate::conversation::message::Message; - use crate::model::ModelConfig; -use crate::providers::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage}; +use crate::providers::base::{ + ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage, +}; use crate::providers::formats::openai::{create_request, get_usage, response_to_message}; use anyhow::Result; use async_trait::async_trait; use rmcp::model::Tool; use serde_json::Value; - pub const XAI_API_HOST: &str = "https://api.x.ai/v1"; pub const XAI_DEFAULT_MODEL: &str = "grok-code-fast-1"; pub const XAI_KNOWN_MODELS: &[&str] = &[ @@ -42,6 +45,7 @@ pub struct XaiProvider { #[serde(skip)] api_client: ApiClient, model: ModelConfig, + supports_streaming: bool, #[serde(skip)] name: String, } @@ -60,13 +64,12 @@ impl XaiProvider { Ok(Self { api_client, model, + supports_streaming: true, name: Self::metadata().name, }) } async fn post(&self, payload: Value) -> Result { - tracing::debug!("xAI request model: {:?}", self.model.model_name); - let response = self .api_client .response_post("chat/completions", &payload) @@ -118,6 +121,7 @@ impl Provider for XaiProvider { messages, tools, &super::utils::ImageFormat::OpenAi, + false, )?; let mut log = RequestLog::start(&self.model, &payload)?; @@ -132,4 +136,40 @@ impl Provider for XaiProvider { log.write(&response, Some(&usage))?; Ok((message, ProviderUsage::new(response_model, usage))) } + + fn supports_streaming(&self) -> bool { + self.supports_streaming + } + + async fn stream( + &self, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result { + let payload = create_request( + &self.model, + system, + messages, + tools, + &super::utils::ImageFormat::OpenAi, + true, + )?; + let mut log = RequestLog::start(&self.model, &payload)?; + + let response = self + .with_retry(|| async { + let resp = self + .api_client + .response_post("chat/completions", &payload) + .await?; + handle_status_openai_compat(resp).await + }) + .await + .inspect_err(|e| { + let _ = log.error(e); + })?; + + stream_openai_compat(response, log) + } } diff --git a/ui/desktop/src/components/settings/app/ExternalBackendSection.tsx b/ui/desktop/src/components/settings/app/ExternalBackendSection.tsx index e16fda9d99..8e52145761 100644 --- a/ui/desktop/src/components/settings/app/ExternalBackendSection.tsx +++ b/ui/desktop/src/components/settings/app/ExternalBackendSection.tsx @@ -100,81 +100,81 @@ export default function ExternalBackendSection() { Goose Server - - By default goose launches a server for you, use this to connect to an external goose - server - - - -
-
-

Use external server

-

- Connect to a goose server running elsewhere (requires app restart) -

-
-
- saveConfig(updateField('enabled', checked))} - disabled={isSaving} - variant="mono" - /> -
-
- - {config.enabled && ( - <> -
- - handleUrlChange(e.target.value)} - onBlur={handleUrlBlur} + + By default goose launches a server for you, use this to connect to an external goose + server + + + +
+
+

Use external server

+

+ Connect to a goose server running elsewhere (requires app restart) +

+
+
+ saveConfig(updateField('enabled', checked))} disabled={isSaving} - className={urlError ? 'border-red-500' : ''} + variant="mono" /> - {urlError && ( -

- - {urlError} +

+
+ + {config.enabled && ( + <> +
+ + handleUrlChange(e.target.value)} + onBlur={handleUrlBlur} + disabled={isSaving} + className={urlError ? 'border-red-500' : ''} + /> + {urlError && ( +

+ + {urlError} +

+ )} +
+ +
+ + updateField('secret', e.target.value)} + onBlur={() => saveConfig(config)} + disabled={isSaving} + /> +

+ The secret key configured on the goosed server (GOOSE_SERVER__SECRET_KEY)

- )} -
+
-
- - updateField('secret', e.target.value)} - onBlur={() => saveConfig(config)} - disabled={isSaving} - /> -

- The secret key configured on the goosed server (GOOSE_SERVER__SECRET_KEY) -

-
- -
-

- Note: Changes require restarting Goose to take effect. New chat - windows will connect to the external server. -

-
- - )} -
-
+
+

+ Note: Changes require restarting Goose to take effect. New chat + windows will connect to the external server. +

+
+ + )} + + ); }