diff --git a/crates/goose-provider-types/src/errors.rs b/crates/goose-provider-types/src/errors.rs index ed38de7e2542..c8d099c29abc 100644 --- a/crates/goose-provider-types/src/errors.rs +++ b/crates/goose-provider-types/src/errors.rs @@ -131,9 +131,20 @@ fn provider_error_from_reqwest(error: &reqwest::Error) -> ProviderError { impl From for ProviderError { fn from(error: anyhow::Error) -> Self { + if let Some(provider_error) = error.downcast_ref::() { + return provider_error.clone(); + } if let Some(reqwest_err) = error.downcast_ref::() { return provider_error_from_reqwest(reqwest_err); } + if error + .downcast_ref::() + .is_some() + { + return ProviderError::NetworkError( + "Request timed out — check your network connection and try again.".to_string(), + ); + } ProviderError::ExecutionError(error.to_string()) } } diff --git a/crates/goose-providers/Cargo.toml b/crates/goose-providers/Cargo.toml index e17e0322f7ce..ab231235b947 100644 --- a/crates/goose-providers/Cargo.toml +++ b/crates/goose-providers/Cargo.toml @@ -58,7 +58,7 @@ include_dir = { workspace = true } [dev-dependencies] test-case = { workspace = true } tempfile = { workspace = true } -tokio = { workspace = true, features = ["rt-multi-thread"] } +tokio = { workspace = true, features = ["io-util", "macros", "net", "rt-multi-thread", "time"] } tokio-stream = { workspace = true } env-lock = { workspace = true } wiremock.workspace = true diff --git a/crates/goose-providers/src/anthropic.rs b/crates/goose-providers/src/anthropic.rs index 09391acdf4f7..8ec61519f4a3 100644 --- a/crates/goose-providers/src/anthropic.rs +++ b/crates/goose-providers/src/anthropic.rs @@ -188,6 +188,7 @@ impl AnthropicProvider { self.api_client .request("v1/messages") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?, ) diff --git a/crates/goose-providers/src/api_client.rs b/crates/goose-providers/src/api_client.rs index b41b35594330..ccb9946c0a91 100644 --- a/crates/goose-providers/src/api_client.rs +++ b/crates/goose-providers/src/api_client.rs @@ -14,7 +14,8 @@ use std::path::PathBuf; use std::sync::Arc; use std::time::Duration; -const DEFAULT_PROVIDER_TIMEOUT_SECS: u64 = 600; +pub const DEFAULT_PROVIDER_TIMEOUT_SECS: u64 = 600; +pub const DEFAULT_CONNECT_TIMEOUT_SECS: u64 = 30; pub type RequestBuilderDecorator = Arc Result + Send + Sync>; @@ -233,6 +234,7 @@ pub struct ApiRequestBuilder<'a> { client: &'a ApiClient, path: &'a str, headers: HeaderMap, + streaming: bool, } impl ApiClient { @@ -255,7 +257,7 @@ impl ApiClient { timeout: Duration, tls_config: Option, ) -> Result { - let mut client_builder = Client::builder().timeout(timeout); + let mut client_builder = Self::client_builder(timeout); if let Some(ref config) = tls_config { client_builder = Self::configure_tls(client_builder, config)?; @@ -283,10 +285,15 @@ impl ApiClient { self.timeout } + fn client_builder(timeout: Duration) -> reqwest::ClientBuilder { + Client::builder() + .connect_timeout(Duration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS)) + .read_timeout(timeout) + } + fn rebuild_client(&mut self) -> Result<()> { - let mut client_builder = Client::builder() - .timeout(self.timeout) - .default_headers(self.default_headers.clone()); + let mut client_builder = + Self::client_builder(self.timeout).default_headers(self.default_headers.clone()); // Configure TLS if needed if let Some(ref tls_config) = self.tls_config { @@ -361,6 +368,7 @@ impl ApiClient { client: self, path, headers: HeaderMap::new(), + streaming: false, } } @@ -434,14 +442,27 @@ impl<'a> ApiRequestBuilder<'a> { } } + pub fn streaming(mut self, streaming: bool) -> Self { + self.streaming = streaming; + self + } + pub async fn api_post(self, payload: &Value) -> Result { let response = self.response_post(payload).await?; ApiResponse::from_response(response).await } + async fn send_bounded(&self, request: reqwest::RequestBuilder) -> Result { + if self.streaming { + Ok(crate::http_status::send_bounded(request, self.client.timeout).await?) + } else { + Ok(request.send().await?) + } + } + pub async fn response_post(self, payload: &Value) -> Result { let request = self.send_request(|url, client| client.post(url)).await?; - Ok(request.json(payload).send().await?) + self.send_bounded(request.json(payload)).await } pub async fn multipart_post(self, form: reqwest::multipart::Form) -> Result { @@ -456,7 +477,7 @@ impl<'a> ApiRequestBuilder<'a> { pub async fn response_get(self) -> Result { let request = self.send_request(|url, client| client.get(url)).await?; - Ok(request.send().await?) + self.send_bounded(request).await } async fn send_request(&self, request_builder: F) -> Result @@ -468,6 +489,10 @@ impl<'a> ApiRequestBuilder<'a> { let mut request = request_builder(url, &self.client.client); request = request.headers(headers); + if !self.streaming { + request = request.timeout(self.client.timeout); + } + if let Some(decorator) = &self.client.request_builder { request = decorator(request)?; } @@ -623,6 +648,198 @@ ShGoCNbfNS+COlPMRAujyDlATZcLs9p4tA== #[cfg(test)] mod tests { use super::*; + use std::net::SocketAddr; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + async fn spawn_chunked_server(gap_ms: u64, chunks: usize) -> SocketAddr { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let Ok((mut sock, _)) = listener.accept().await else { + break; + }; + tokio::spawn(async move { + let mut buf = [0u8; 8192]; + let _ = sock.read(&mut buf).await; + if sock + .write_all( + b"HTTP/1.1 200 OK\r\n\ + content-type: text/event-stream\r\n\ + transfer-encoding: chunked\r\n\r\n", + ) + .await + .is_err() + { + return; + } + for i in 0..chunks { + if i > 0 { + tokio::time::sleep(Duration::from_millis(gap_ms)).await; + } + let data = format!("data: {}\n\n", i); + let chunk = format!("{:x}\r\n{}\r\n", data.len(), data); + if sock.write_all(chunk.as_bytes()).await.is_err() { + return; + } + let _ = sock.flush().await; + } + let _ = sock.write_all(b"0\r\n\r\n").await; + }); + } + }); + addr + } + + fn client_with_timeout(addr: SocketAddr, timeout_ms: u64) -> ApiClient { + let mut client = ApiClient::with_timeout_and_tls( + format!("http://{}", addr), + AuthMethod::NoAuth, + Duration::from_millis(timeout_ms), + None, + ) + .unwrap(); + client.client = Client::builder() + .no_proxy() + .connect_timeout(Duration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS)) + .read_timeout(client.timeout) + .build() + .unwrap(); + client + } + + async fn drain_counting_data_lines(mut response: Response) -> Result { + let mut body = Vec::new(); + while let Some(chunk) = response.chunk().await? { + body.extend_from_slice(&chunk); + } + Ok(String::from_utf8_lossy(&body).matches("data:").count()) + } + + #[tokio::test] + async fn streaming_request_survives_beyond_total_timeout() { + let addr = spawn_chunked_server(50, 12).await; + let client = client_with_timeout(addr, 400); + + let response = client + .request("v1/messages") + .streaming(true) + .response_post(&serde_json::json!({})) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let count = drain_counting_data_lines(response).await.unwrap(); + assert_eq!(count, 12); + } + + #[tokio::test] + async fn streaming_request_fails_when_stream_stalls() { + let addr = spawn_chunked_server(5_000, 2).await; + let client = client_with_timeout(addr, 400); + + let response = client + .request("v1/messages") + .streaming(true) + .response_post(&serde_json::json!({})) + .await + .unwrap(); + + let err = drain_counting_data_lines(response) + .await + .expect_err("stalled stream should time out, not complete"); + assert!(err.is_timeout(), "expected a timeout error, got: {err}"); + } + + #[tokio::test] + async fn non_streaming_request_enforces_total_deadline() { + let addr = spawn_chunked_server(50, 12).await; + let client = client_with_timeout(addr, 400); + + let response = client + .request("v1/messages") + .response_post(&serde_json::json!({})) + .await + .unwrap(); + + let err = drain_counting_data_lines(response) + .await + .expect_err("total deadline should cut off the response body"); + assert!(err.is_timeout(), "expected a timeout error, got: {err}"); + } + + #[tokio::test] + async fn streaming_request_times_out_before_response_headers() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let Ok((mut sock, _)) = listener.accept().await else { + break; + }; + tokio::spawn(async move { + let mut buf = [0u8; 8192]; + while sock.read(&mut buf).await.is_ok_and(|n| n > 0) {} + }); + } + }); + let client = client_with_timeout(addr, 400); + + let started = std::time::Instant::now(); + let err = client + .request("v1/messages") + .streaming(true) + .response_post(&serde_json::json!({})) + .await + .expect_err("the phase before the response body must stay bounded"); + assert!( + started.elapsed() < Duration::from_secs(5), + "should fail near the configured timeout, took {:?}", + started.elapsed() + ); + assert!(matches!( + crate::errors::ProviderError::from(err), + crate::errors::ProviderError::NetworkError(message) + if message.starts_with("Request timed out") + )); + } + + #[tokio::test] + async fn streaming_error_body_shares_send_deadline() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut buf = [0u8; 8192]; + let _ = socket.read(&mut buf).await; + tokio::time::sleep(Duration::from_millis(300)).await; + socket + .write_all(b"HTTP/1.1 500 Internal Server Error\r\ncontent-length: 1\r\n\r\n") + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(300)).await; + let _ = socket.write_all(b"x").await; + }); + + let client = client_with_timeout(addr, 400); + let started = std::time::Instant::now(); + let response = client + .request("v1/messages") + .streaming(true) + .response_post(&serde_json::json!({})) + .await + .unwrap(); + crate::http_status::handle_status(response) + .await + .unwrap_err(); + + assert!( + started.elapsed() < Duration::from_millis(550), + "send and error body used separate deadlines: {:?}", + started.elapsed() + ); + } #[test] fn test_model_headers_applied_and_override_static_headers() { diff --git a/crates/goose-providers/src/databricks.rs b/crates/goose-providers/src/databricks.rs index ad95803b49ce..eb8a4380aa56 100644 --- a/crates/goose-providers/src/databricks.rs +++ b/crates/goose-providers/src/databricks.rs @@ -601,6 +601,7 @@ impl Provider for DatabricksProvider { .api_client .request(&path) .model_headers(model_config)? + .streaming(true) .response_post(&payload_clone) .await?; handle_status(resp).await @@ -666,12 +667,15 @@ impl Provider for DatabricksProvider { .api_client .request(&path) .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; if !resp.status().is_success() { let status = resp.status(); let url = sanitize_url(resp.url().as_str()); - let error_text = resp.text().await.unwrap_or_default(); + let error_text = crate::http_status::read_error_body(resp) + .await + .unwrap_or_default(); let json_payload = serde_json::from_str::(&error_text).ok(); return Err(map_http_error_to_provider_error(status, json_payload, &url)); @@ -688,12 +692,15 @@ impl Provider for DatabricksProvider { .api_client .request(&path) .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; if !resp.status().is_success() { let status = resp.status(); let url = sanitize_url(resp.url().as_str()); - let error_text = resp.text().await.unwrap_or_default(); + let error_text = crate::http_status::read_error_body(resp) + .await + .unwrap_or_default(); let json_payload = serde_json::from_str::(&error_text).ok(); return Err(map_http_error_to_provider_error( status, diff --git a/crates/goose-providers/src/databricks_v2.rs b/crates/goose-providers/src/databricks_v2.rs index 8a2c1f8c431f..0557676527a8 100644 --- a/crates/goose-providers/src/databricks_v2.rs +++ b/crates/goose-providers/src/databricks_v2.rs @@ -207,6 +207,7 @@ impl DatabricksV2Provider { .api_client .request("ai-gateway/openai/v1/responses") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; handle_status(resp).await @@ -245,6 +246,7 @@ impl DatabricksV2Provider { .api_client .request("ai-gateway/mlflow/v1/chat/completions") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; handle_status(resp).await @@ -281,6 +283,7 @@ impl DatabricksV2Provider { .api_client .request("ai-gateway/anthropic/v1/messages") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; handle_status(resp).await diff --git a/crates/goose-providers/src/google.rs b/crates/goose-providers/src/google.rs index 16599fd122b5..bc9e7eab7e1e 100644 --- a/crates/goose-providers/src/google.rs +++ b/crates/goose-providers/src/google.rs @@ -107,6 +107,7 @@ impl GoogleProvider { .api_client .request(&path) .model_headers(model_config)? + .streaming(true) .response_post(payload) .await?; handle_status(response).await diff --git a/crates/goose-providers/src/http_status.rs b/crates/goose-providers/src/http_status.rs index 13147ccadd4a..51f8e8e4ff0a 100644 --- a/crates/goose-providers/src/http_status.rs +++ b/crates/goose-providers/src/http_status.rs @@ -248,12 +248,45 @@ pub fn map_http_error_to_provider_error( error } +#[derive(Clone, Copy)] +pub struct ResponseDeadline(tokio::time::Instant); + +pub fn set_response_deadline(response: &mut Response, deadline: tokio::time::Instant) { + response.extensions_mut().insert(ResponseDeadline(deadline)); +} + +pub async fn send_bounded( + request: reqwest::RequestBuilder, + timeout: Duration, +) -> Result { + let deadline = tokio::time::Instant::now() + timeout; + let mut response = tokio::time::timeout_at(deadline, request.send()) + .await + .map_err(|_| { + ProviderError::NetworkError( + "Request timed out — check your network connection and try again.".to_string(), + ) + })??; + set_response_deadline(&mut response, deadline); + Ok(response) +} + +pub async fn read_error_body(response: Response) -> Option { + match response.extensions().get::().copied() { + Some(ResponseDeadline(deadline)) => tokio::time::timeout_at(deadline, response.text()) + .await + .ok() + .and_then(Result::ok), + None => response.text().await.ok(), + } +} + pub async fn handle_status(response: Response) -> Result { let status = response.status(); if !status.is_success() { let url = sanitize_url(response.url().as_str()); let headers = response.headers().clone(); - let body = response.text().await.unwrap_or_default(); + let body = read_error_body(response).await.unwrap_or_default(); let payload = serde_json::from_str::(&body).ok(); let mut err = map_http_error_to_provider_error(status, payload.clone(), &url); if let ProviderError::RateLimitExceeded { details, .. } = &err { diff --git a/crates/goose-providers/src/ollama.rs b/crates/goose-providers/src/ollama.rs index 2a4cdbab8486..c1de5a336370 100644 --- a/crates/goose-providers/src/ollama.rs +++ b/crates/goose-providers/src/ollama.rs @@ -428,6 +428,7 @@ impl Provider for OllamaProvider { .api_client .request("v1/chat/completions") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; handle_status(resp).await diff --git a/crates/goose-providers/src/openai.rs b/crates/goose-providers/src/openai.rs index 15a5c983295f..b256ca989da8 100644 --- a/crates/goose-providers/src/openai.rs +++ b/crates/goose-providers/src/openai.rs @@ -311,6 +311,7 @@ impl OpenAiProvider { OPEN_AI_DEFAULT_RESPONSES_PATH, )) .model_headers(model_config)? + .streaming(self.supports_streaming) .response_post(&payload) .await?, ) @@ -797,6 +798,7 @@ impl Provider for OpenAiProvider { .api_client .request(&self.base_path) .model_headers(model_config)? + .streaming(self.supports_streaming) .response_post(&payload) .await?; handle_status(resp).await diff --git a/crates/goose-providers/src/openai_compatible.rs b/crates/goose-providers/src/openai_compatible.rs index d7f752bde616..f40d756da26c 100644 --- a/crates/goose-providers/src/openai_compatible.rs +++ b/crates/goose-providers/src/openai_compatible.rs @@ -112,6 +112,7 @@ impl OpenAiCompatibleProvider { self.api_client .request(&path) .model_headers(model_config)? + .streaming(self.supports_streaming) .response_post(&payload) .await?, ) diff --git a/crates/goose-providers/src/snowflake.rs b/crates/goose-providers/src/snowflake.rs index d30cb26d9433..d0fa2bc06566 100644 --- a/crates/goose-providers/src/snowflake.rs +++ b/crates/goose-providers/src/snowflake.rs @@ -100,12 +100,25 @@ impl SnowflakeProvider { .api_client .request("api/v2/cortex/inference:complete") .model_headers(model_config)? + .streaming(true) .response_post(payload) .await?; let status = response.status(); let url = sanitize_url(response.url().as_str()); - let payload_text: String = response.text().await.ok().unwrap_or_default(); + let is_json = response + .headers() + .get(reqwest::header::CONTENT_TYPE) + .and_then(|v| v.to_str().ok()) + .map(|v| v.to_ascii_lowercase()) + .is_some_and(|v| v.contains("json")); + let payload_text: String = if status.is_success() && !is_json { + response.text().await.ok().unwrap_or_default() + } else { + crate::http_status::read_error_body(response) + .await + .unwrap_or_default() + }; if status.is_success() { if let Ok(payload) = serde_json::from_str::(&payload_text) { diff --git a/crates/goose/src/providers/anthropic_def.rs b/crates/goose/src/providers/anthropic_def.rs index 6388511a49bc..2efffb277374 100644 --- a/crates/goose/src/providers/anthropic_def.rs +++ b/crates/goose/src/providers/anthropic_def.rs @@ -39,14 +39,23 @@ async fn from_env( .get_param("ANTHROPIC_HOST") .unwrap_or_else(|_| "https://api.anthropic.com".to_string()); + let timeout_secs: u64 = config + .get_param("ANTHROPIC_TIMEOUT") + .unwrap_or(crate::providers::base::DEFAULT_PROVIDER_TIMEOUT_SECS); + let auth = AuthMethod::ApiKey { header_name: "x-api-key".to_string(), key: api_key, }; - let api_client = ApiClient::new_with_tls(host, auth, tls_config)? - .with_request_builder(crate::session_context::session_id_request_builder()) - .with_header("anthropic-version", ANTHROPIC_API_VERSION)?; + let api_client = ApiClient::with_timeout_and_tls( + host, + auth, + std::time::Duration::from_secs(timeout_secs), + tls_config, + )? + .with_request_builder(crate::session_context::session_id_request_builder()) + .with_header("anthropic-version", ANTHROPIC_API_VERSION)?; Ok(AnthropicProviderBuilder::new(api_client).build()) } diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index 8baee4838874..5f6ce506fdfb 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -6,7 +6,9 @@ pub use goose_providers::conversation::token_usage::{ }; use serde::{Deserialize, Serialize}; -pub const DEFAULT_PROVIDER_TIMEOUT_SECS: u64 = 600; +pub use goose_providers::api_client::{ + DEFAULT_CONNECT_TIMEOUT_SECS, DEFAULT_PROVIDER_TIMEOUT_SECS, +}; use crate::config::ExtensionConfig; diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index d6fee4ceceee..7b48823baaa8 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -1,6 +1,9 @@ use std::collections::HashMap; -use super::base::{ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata}; +use super::base::{ + ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, + DEFAULT_CONNECT_TIMEOUT_SECS, DEFAULT_PROVIDER_TIMEOUT_SECS, +}; use super::openai_compatible::{handle_status, stream_responses_compat}; use super::retry::{ProviderRetry, RetryConfig}; use crate::conversation::message::Message; @@ -177,7 +180,12 @@ impl BedrockProvider { name: BEDROCK_PROVIDER_NAME.to_string(), region: resolved_region, bearer_token, - http_client: reqwest::Client::new(), + http_client: reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS)) + .read_timeout(std::time::Duration::from_secs( + DEFAULT_PROVIDER_TIMEOUT_SECS, + )) + .build()?, mantle_base_url: None, }) } @@ -270,10 +278,11 @@ impl BedrockProvider { } } - let response = req - .send() - .await - .map_err(|e| ProviderError::RequestFailed(format!("Mantle request failed: {}", e)))?; + let response = goose_providers::http_status::send_bounded( + req, + std::time::Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS), + ) + .await?; handle_status(response).await } diff --git a/crates/goose/src/providers/chatgpt_codex.rs b/crates/goose/src/providers/chatgpt_codex.rs index d683a9eef40b..3622898cd2f0 100644 --- a/crates/goose/src/providers/chatgpt_codex.rs +++ b/crates/goose/src/providers/chatgpt_codex.rs @@ -1,7 +1,10 @@ use crate::config::paths::Paths; use crate::conversation::message::{Message, MessageContent}; use crate::providers::api_client::{AuthProvider, RequestBuilderDecorator}; -use crate::providers::base::{ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata}; +use crate::providers::base::{ + ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, + DEFAULT_CONNECT_TIMEOUT_SECS, DEFAULT_PROVIDER_TIMEOUT_SECS, +}; use crate::providers::openai_compatible::handle_status; use crate::providers::private_file::write_private_file; use crate::providers::retry::ProviderRetry; @@ -921,7 +924,13 @@ impl ChatGptCodexProvider { ); } - let client = reqwest::Client::new(); + let client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS)) + .read_timeout(std::time::Duration::from_secs( + DEFAULT_PROVIDER_TIMEOUT_SECS, + )) + .build() + .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; let request = client .post(format!("{}/responses", CODEX_API_ENDPOINT)) .header( @@ -932,11 +941,13 @@ impl ChatGptCodexProvider { .headers(headers) .json(payload); - let response = (self.request_builder)(request) - .map_err(|e| ProviderError::ExecutionError(e.to_string()))? - .send() - .await - .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; + let request = (self.request_builder)(request) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; + let response = goose_providers::http_status::send_bounded( + request, + std::time::Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS), + ) + .await?; handle_status(response).await } diff --git a/crates/goose/src/providers/gcpvertexai.rs b/crates/goose/src/providers/gcpvertexai.rs index 6a2ea2ebc47f..07f1aa218396 100644 --- a/crates/goose/src/providers/gcpvertexai.rs +++ b/crates/goose/src/providers/gcpvertexai.rs @@ -18,7 +18,7 @@ use crate::conversation::message::Message; use crate::providers::api_client::RequestBuilderDecorator; use crate::providers::base::{ ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, - DEFAULT_PROVIDER_TIMEOUT_SECS, + DEFAULT_CONNECT_TIMEOUT_SECS, DEFAULT_PROVIDER_TIMEOUT_SECS, }; use goose_providers::model::ModelConfig; @@ -30,6 +30,7 @@ use crate::providers::gcpauth::GcpAuth; use crate::providers::openai_compatible::{map_http_error_to_provider_error, sanitize_url}; use crate::providers::retry::RetryConfig; use goose_providers::errors::ProviderError; +use goose_providers::http_status::read_error_body; use goose_providers::request_log::{start_log, LoggerHandleExt}; use rmcp::model::Tool; @@ -174,7 +175,8 @@ impl GcpVertexAIProvider { let host = Self::build_host_url(&location); let client = Client::builder() - .timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) + .connect_timeout(Duration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS)) + .read_timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .build()?; let auth = GcpAuth::new().await?; @@ -327,11 +329,13 @@ impl GcpVertexAIProvider { } } - let response = (self.request_builder)(request) - .map_err(|e| ProviderError::ExecutionError(e.to_string()))? - .send() - .await - .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; + let request = (self.request_builder)(request) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; + let response = goose_providers::http_status::send_bounded( + request, + Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS), + ) + .await?; let status = response.status(); @@ -345,7 +349,8 @@ impl GcpVertexAIProvider { }), ); } - let msg = rate_limit_error_message(&response.text().await.unwrap_or_default()); + let msg = + rate_limit_error_message(&read_error_body(response).await.unwrap_or_default()); tracing::warn!("429 (attempt {rate_limit_attempts}/{max_retries}): {msg}"); last_error = Some(ProviderError::RateLimitExceeded { details: msg, @@ -389,7 +394,7 @@ impl GcpVertexAIProvider { ))); } else { let url = sanitize_url(response.url().as_str()); - let response_text = response.text().await.unwrap_or_default(); + let response_text = read_error_body(response).await.unwrap_or_default(); let payload = serde_json::from_str::(&response_text).ok(); return Err(map_http_error_to_provider_error(status, payload, &url)); } @@ -459,6 +464,7 @@ impl GcpVertexAIProvider { let response = match self .client .post(&url) + .timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .header("Authorization", &auth_header) .json(&payload) .send() diff --git a/crates/goose/src/providers/gemini_oauth.rs b/crates/goose/src/providers/gemini_oauth.rs index a954838e48a5..39bdb7f437bb 100644 --- a/crates/goose/src/providers/gemini_oauth.rs +++ b/crates/goose/src/providers/gemini_oauth.rs @@ -3,7 +3,7 @@ use crate::conversation::message::Message; use crate::providers::api_client::RequestBuilderDecorator; use crate::providers::base::{ ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, - DEFAULT_PROVIDER_TIMEOUT_SECS, + DEFAULT_CONNECT_TIMEOUT_SECS, DEFAULT_PROVIDER_TIMEOUT_SECS, }; use crate::providers::formats::google::{create_request, response_to_streaming_message}; use crate::providers::google::GOOGLE_DOC_URL; @@ -40,7 +40,8 @@ use tokio_util::io::StreamReader; static HTTP_CLIENT: LazyLock = LazyLock::new(|| { reqwest::Client::builder() - .timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) + .connect_timeout(Duration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS)) + .read_timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .build() .expect("failed to build HTTP client") }); @@ -254,6 +255,7 @@ async fn exchange_code_for_tokens( let resp = client .post(GOOGLE_TOKEN_ENDPOINT) + .timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .header("Content-Type", "application/x-www-form-urlencoded") .form(¶ms) .send() @@ -281,6 +283,7 @@ async fn refresh_access_token(refresh_token: &str) -> Result { let resp = client .post(GOOGLE_TOKEN_ENDPOINT) + .timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .header("Content-Type", "application/x-www-form-urlencoded") .form(¶ms) .send() @@ -343,6 +346,7 @@ async fn code_assist_request(access_token: &str, method: &str, body: &Value) -> let client = &*HTTP_CLIENT; let resp = client .post(&url) + .timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .header("Authorization", format!("Bearer {}", access_token)) .header("Content-Type", "application/json") .json(body) @@ -371,6 +375,7 @@ async fn code_assist_get(access_token: &str, path: &str) -> Result { let client = &*HTTP_CLIENT; let resp = client .get(&url) + .timeout(Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .header("Authorization", format!("Bearer {}", access_token)) .send() .await?; @@ -884,18 +889,19 @@ impl GeminiOAuthProvider { ) .header("Content-Type", "application/json"); - let response = (self.request_builder)(request.json(&wrapped)) - .map_err(|e| ProviderError::ExecutionError(e.to_string()))? - .send() - .await - .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; + let request = (self.request_builder)(request.json(&wrapped)) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; + let response = goose_providers::http_status::send_bounded( + request, + Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS), + ) + .await?; if !response.status().is_success() { let status = response.status(); - let text = response - .text() + let text = goose_providers::http_status::read_error_body(response) .await - .unwrap_or_else(|_| "unknown error".to_string()); + .unwrap_or_else(|| "unknown error".to_string()); if status == reqwest::StatusCode::TOO_MANY_REQUESTS { // Parse retry delay from the error message if available diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index e99b7b7034d8..d1550afcae9b 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -265,6 +265,7 @@ impl GithubCopilotProvider { is_user_initiated: bool, payload: &mut Value, has_images: bool, + streaming: bool, ) -> Result { let (endpoint, token) = self.get_api_info().await?; let auth = AuthMethod::BearerToken(token); @@ -281,6 +282,7 @@ impl GithubCopilotProvider { api_client .request(path) .model_headers(model_config)? + .streaming(streaming) .response_post(payload) .await .map_err(|e| e.into()) @@ -420,6 +422,7 @@ impl GithubCopilotProvider { is_user_initiated, &mut payload_clone, has_images, + true, ) .await?; handle_status(resp).await @@ -467,6 +470,7 @@ impl GithubCopilotProvider { is_user_initiated, &mut payload_clone, has_images, + true, ) .await?; handle_status(resp).await @@ -497,6 +501,7 @@ impl GithubCopilotProvider { is_user_initiated, &mut payload_clone, has_images, + false, ) .await }) diff --git a/crates/goose/src/providers/kimicode.rs b/crates/goose/src/providers/kimicode.rs index 5568ef9fd480..3e0a38e39194 100644 --- a/crates/goose/src/providers/kimicode.rs +++ b/crates/goose/src/providers/kimicode.rs @@ -19,7 +19,7 @@ use uuid::Uuid; use super::api_client::RequestBuilderDecorator; use super::base::{ ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, - DEFAULT_PROVIDER_TIMEOUT_SECS, + DEFAULT_CONNECT_TIMEOUT_SECS, DEFAULT_PROVIDER_TIMEOUT_SECS, }; use super::formats::anthropic::{create_request, response_to_streaming_message}; use super::oauth_device_flow::{ @@ -172,7 +172,8 @@ impl KimiCodeProvider { _tls_config: Option, ) -> Result { let client = Client::builder() - .timeout(StdDuration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) + .connect_timeout(StdDuration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS)) + .read_timeout(StdDuration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .build()?; let device_id = Self::get_or_create_device_id().await?; Ok(Self { @@ -326,11 +327,13 @@ impl KimiCodeProvider { .headers(self.kimi_headers()) .json(payload); - (self.request_builder)(builder) - .map_err(|e| ProviderError::ExecutionError(e.to_string()))? - .send() - .await - .map_err(|e| ProviderError::RequestFailed(e.to_string())) + let request = (self.request_builder)(builder) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; + goose_providers::http_status::send_bounded( + request, + StdDuration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS), + ) + .await } } @@ -454,6 +457,7 @@ impl Provider for KimiCodeProvider { let resp = self .client .get(format!("{}/v1/models", self.api_base)) + .timeout(StdDuration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .bearer_auth(access_token) .headers(self.kimi_headers()) .send() diff --git a/crates/goose/src/providers/nanogpt.rs b/crates/goose/src/providers/nanogpt.rs index df7fcb57fe57..2ced5d7b2893 100644 --- a/crates/goose/src/providers/nanogpt.rs +++ b/crates/goose/src/providers/nanogpt.rs @@ -198,6 +198,7 @@ impl Provider for NanoGptProvider { .api_client .request("chat/completions") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; handle_status(resp).await diff --git a/crates/goose/src/providers/oauth_device_flow.rs b/crates/goose/src/providers/oauth_device_flow.rs index d257d67b2df8..6baff9a54108 100644 --- a/crates/goose/src/providers/oauth_device_flow.rs +++ b/crates/goose/src/providers/oauth_device_flow.rs @@ -287,7 +287,12 @@ async fn send_request( url: &str, body: &T, ) -> reqwest::Result { - let builder = client.post(url).headers(cfg.extra_headers.clone()); + let builder = client + .post(url) + .timeout(std::time::Duration::from_secs( + super::base::DEFAULT_PROVIDER_TIMEOUT_SECS, + )) + .headers(cfg.extra_headers.clone()); let builder = match cfg.encoding { RequestEncoding::Form => builder.form(body), RequestEncoding::Json => builder.json(body), diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index 631745ccfe49..177a6356d2e8 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -350,6 +350,7 @@ impl Provider for OpenRouterProvider { .api_client .request("api/v1/chat/completions") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; handle_status(resp).await diff --git a/crates/goose/src/providers/tetrate.rs b/crates/goose/src/providers/tetrate.rs index 2bb64231d16b..0524a10e3822 100644 --- a/crates/goose/src/providers/tetrate.rs +++ b/crates/goose/src/providers/tetrate.rs @@ -154,6 +154,7 @@ impl Provider for TetrateProvider { .api_client .request("v1/chat/completions") .model_headers(model_config)? + .streaming(true) .response_post(&payload) .await?; let resp = handle_status(resp) @@ -168,22 +169,21 @@ impl Provider for TetrateProvider { .is_some_and(|v| v.contains("json")); if is_json { - // Streaming responses should be SSE; when we get JSON instead, parse it to map - // explicit error payloads and otherwise fail as a protocol mismatch. - let body = handle_response_openai_compat(resp) + let body = goose_providers::http_status::read_error_body(resp) .await - .map_err(Self::enrich_credits_error)?; - if body.get("error").is_some() { - return Err(Self::error_from_tetrate_error_payload( - body, - "v1/chat/completions", - )); + .unwrap_or_default(); + if let Ok(payload) = serde_json::from_str::(&body) { + if payload.get("error").is_some() { + return Err(Self::error_from_tetrate_error_payload( + payload, + "v1/chat/completions", + )); + } } - return Err(ProviderError::ExecutionError( - "Expected streaming response but received non-streaming payload" - .to_string(), - )); + return Err(ProviderError::ExecutionError(format!( + "Expected streaming response but received non-streaming payload: {body}" + ))); } Ok(resp)