fix(providers): stop killing streaming responses at the total request timeout (#10620)
Co-authored-by: Douwe M Osinga <douwe@sidewalklabs.com>
This commit is contained in:
@@ -131,9 +131,20 @@ fn provider_error_from_reqwest(error: &reqwest::Error) -> ProviderError {
|
||||
|
||||
impl From<anyhow::Error> for ProviderError {
|
||||
fn from(error: anyhow::Error) -> Self {
|
||||
if let Some(provider_error) = error.downcast_ref::<ProviderError>() {
|
||||
return provider_error.clone();
|
||||
}
|
||||
if let Some(reqwest_err) = error.downcast_ref::<reqwest::Error>() {
|
||||
return provider_error_from_reqwest(reqwest_err);
|
||||
}
|
||||
if error
|
||||
.downcast_ref::<tokio::time::error::Elapsed>()
|
||||
.is_some()
|
||||
{
|
||||
return ProviderError::NetworkError(
|
||||
"Request timed out — check your network connection and try again.".to_string(),
|
||||
);
|
||||
}
|
||||
ProviderError::ExecutionError(error.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -188,6 +188,7 @@ impl AnthropicProvider {
|
||||
self.api_client
|
||||
.request("v1/messages")
|
||||
.model_headers(model_config)?
|
||||
.streaming(true)
|
||||
.response_post(&payload)
|
||||
.await?,
|
||||
)
|
||||
|
||||
@@ -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<dyn Fn(reqwest::RequestBuilder) -> Result<reqwest::RequestBuilder> + 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<TlsConfig>,
|
||||
) -> Result<Self> {
|
||||
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<ApiResponse> {
|
||||
let response = self.response_post(payload).await?;
|
||||
ApiResponse::from_response(response).await
|
||||
}
|
||||
|
||||
async fn send_bounded(&self, request: reqwest::RequestBuilder) -> Result<Response> {
|
||||
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<Response> {
|
||||
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<Response> {
|
||||
@@ -456,7 +477,7 @@ impl<'a> ApiRequestBuilder<'a> {
|
||||
|
||||
pub async fn response_get(self) -> Result<Response> {
|
||||
let request = self.send_request(|url, client| client.get(url)).await?;
|
||||
Ok(request.send().await?)
|
||||
self.send_bounded(request).await
|
||||
}
|
||||
|
||||
async fn send_request<F>(&self, request_builder: F) -> Result<reqwest::RequestBuilder>
|
||||
@@ -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<usize, reqwest::Error> {
|
||||
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() {
|
||||
|
||||
@@ -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::<Value>(&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::<Value>(&error_text).ok();
|
||||
return Err(map_http_error_to_provider_error(
|
||||
status,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -107,6 +107,7 @@ impl GoogleProvider {
|
||||
.api_client
|
||||
.request(&path)
|
||||
.model_headers(model_config)?
|
||||
.streaming(true)
|
||||
.response_post(payload)
|
||||
.await?;
|
||||
handle_status(response).await
|
||||
|
||||
@@ -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<Response, ProviderError> {
|
||||
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<String> {
|
||||
match response.extensions().get::<ResponseDeadline>().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<Response, ProviderError> {
|
||||
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::<Value>(&body).ok();
|
||||
let mut err = map_http_error_to_provider_error(status, payload.clone(), &url);
|
||||
if let ProviderError::RateLimitExceeded { details, .. } = &err {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -112,6 +112,7 @@ impl OpenAiCompatibleProvider {
|
||||
self.api_client
|
||||
.request(&path)
|
||||
.model_headers(model_config)?
|
||||
.streaming(self.supports_streaming)
|
||||
.response_post(&payload)
|
||||
.await?,
|
||||
)
|
||||
|
||||
@@ -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::<Value>(&payload_text) {
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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::<Value>(&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()
|
||||
|
||||
@@ -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<reqwest::Client> = 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<TokenResponse> {
|
||||
|
||||
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<Value> {
|
||||
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
|
||||
|
||||
@@ -265,6 +265,7 @@ impl GithubCopilotProvider {
|
||||
is_user_initiated: bool,
|
||||
payload: &mut Value,
|
||||
has_images: bool,
|
||||
streaming: bool,
|
||||
) -> Result<Response, ProviderError> {
|
||||
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
|
||||
})
|
||||
|
||||
@@ -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<crate::providers::api_client::TlsConfig>,
|
||||
) -> Result<Self> {
|
||||
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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -287,7 +287,12 @@ async fn send_request<T: Serialize + ?Sized>(
|
||||
url: &str,
|
||||
body: &T,
|
||||
) -> reqwest::Result<reqwest::Response> {
|
||||
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),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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::<serde_json::Value>(&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)
|
||||
|
||||
Reference in New Issue
Block a user