Mark stream decode errors retryable (#9723)
Signed-off-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -51,6 +51,10 @@ pub enum ProviderError {
|
||||
}
|
||||
|
||||
impl ProviderError {
|
||||
pub fn stream_decode_error(error: impl std::fmt::Display) -> Self {
|
||||
ProviderError::NetworkError(format!("Stream decode error: {error}"))
|
||||
}
|
||||
|
||||
pub fn telemetry_type(&self) -> &'static str {
|
||||
match self {
|
||||
ProviderError::Authentication(_) => "auth",
|
||||
@@ -73,11 +77,12 @@ impl ProviderError {
|
||||
}
|
||||
|
||||
/// Recover a typed `ProviderError` from a streaming decode error, falling
|
||||
/// back to `RequestFailed` for errors that did not originate as one.
|
||||
/// back to a retryable stream decode error for errors that did not
|
||||
/// originate as one.
|
||||
pub fn from_stream_error(error: anyhow::Error) -> Self {
|
||||
error
|
||||
.downcast()
|
||||
.unwrap_or_else(|e| ProviderError::RequestFailed(format!("Stream decode error: {e}")))
|
||||
.unwrap_or_else(ProviderError::stream_decode_error)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -866,7 +866,9 @@ fn strip_data_prefix(line: &str) -> Option<&str> {
|
||||
|
||||
fn parse_streaming_chunk(line: &str) -> Result<StreamingChunk, ProviderError> {
|
||||
let value: Value = serde_json::from_str(line).map_err(|e| {
|
||||
ProviderError::RequestFailed(format!("Failed to parse streaming chunk: {e}: {line:?}"))
|
||||
ProviderError::stream_decode_error(format!(
|
||||
"Failed to parse streaming chunk: {e}: {line:?}"
|
||||
))
|
||||
})?;
|
||||
|
||||
if let Some(error) = value.get("error") {
|
||||
@@ -886,7 +888,9 @@ fn parse_streaming_chunk(line: &str) -> Result<StreamingChunk, ProviderError> {
|
||||
}
|
||||
|
||||
serde_json::from_value(value).map_err(|e| {
|
||||
ProviderError::RequestFailed(format!("Failed to parse streaming chunk: {e}: {line:?}"))
|
||||
ProviderError::stream_decode_error(format!(
|
||||
"Failed to parse streaming chunk: {e}: {line:?}"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -1020,7 +1020,10 @@ impl Provider for ChatGptCodexProvider {
|
||||
let message_stream = responses_api_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)))?;
|
||||
let (message, usage) = message.map_err(|e| {
|
||||
e.downcast::<ProviderError>()
|
||||
.unwrap_or_else(ProviderError::stream_decode_error)
|
||||
})?;
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
|
||||
@@ -433,7 +433,9 @@ where
|
||||
.get("status")
|
||||
.and_then(|s| s.as_str())
|
||||
.unwrap_or("UNKNOWN");
|
||||
Err(anyhow::anyhow!("Google API error ({}): {}", status, message))?;
|
||||
Err::<(), ProviderError>(ProviderError::RequestFailed(format!(
|
||||
"Google API error ({status}): {message}"
|
||||
)))?;
|
||||
}
|
||||
|
||||
if let Ok(usage) = get_usage(&chunk) {
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::mcp_utils::extract_text_from_resource;
|
||||
use crate::model::ModelConfig;
|
||||
use anyhow::{anyhow, Error};
|
||||
use anyhow::Error;
|
||||
use async_stream::try_stream;
|
||||
use chrono;
|
||||
use futures::Stream;
|
||||
use goose_providers::conversation::token_usage::{ProviderUsage, Usage};
|
||||
use goose_providers::errors::ProviderError;
|
||||
use goose_providers::formats::openai::{
|
||||
extract_reasoning_effort, is_openai_responses_model, openai_reasoning_effort_for_thinking,
|
||||
};
|
||||
@@ -249,11 +250,10 @@ fn is_known_responses_stream_event_type(event_type: &str) -> bool {
|
||||
|
||||
fn parse_responses_stream_event(data_line: &str) -> anyhow::Result<Option<ResponsesStreamEvent>> {
|
||||
let raw_event: Value = serde_json::from_str(data_line).map_err(|e| {
|
||||
anyhow!(
|
||||
ProviderError::stream_decode_error(format!(
|
||||
"Failed to parse Responses stream event: {}: {:?}",
|
||||
e,
|
||||
data_line
|
||||
)
|
||||
e, data_line
|
||||
))
|
||||
})?;
|
||||
|
||||
let Some(event_type) = raw_event.get("type").and_then(Value::as_str) else {
|
||||
@@ -265,11 +265,10 @@ fn parse_responses_stream_event(data_line: &str) -> anyhow::Result<Option<Respon
|
||||
}
|
||||
|
||||
let event = serde_json::from_value(raw_event).map_err(|e| {
|
||||
anyhow!(
|
||||
ProviderError::stream_decode_error(format!(
|
||||
"Failed to parse Responses stream event: {}: {:?}",
|
||||
e,
|
||||
data_line
|
||||
)
|
||||
e, data_line
|
||||
))
|
||||
})?;
|
||||
Ok(Some(event))
|
||||
}
|
||||
@@ -911,11 +910,17 @@ where
|
||||
}
|
||||
|
||||
ResponsesStreamEvent::ResponseFailed { error, .. } => {
|
||||
Err(anyhow!("Responses API failed: {:?}", error))?;
|
||||
Err::<(), ProviderError>(ProviderError::RequestFailed(format!(
|
||||
"Responses API failed: {:?}",
|
||||
error
|
||||
)))?;
|
||||
}
|
||||
|
||||
ResponsesStreamEvent::Error { error } => {
|
||||
Err(anyhow!("Responses API error: {:?}", error))?;
|
||||
Err::<(), ProviderError>(ProviderError::RequestFailed(format!(
|
||||
"Responses API error: {:?}",
|
||||
error
|
||||
)))?;
|
||||
}
|
||||
|
||||
_ => {
|
||||
|
||||
@@ -1015,9 +1015,10 @@ impl Provider for GeminiOAuthProvider {
|
||||
let message_stream = response_to_streaming_message(raw_lines);
|
||||
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))
|
||||
)?;
|
||||
let (message, usage) = message.map_err(|e| {
|
||||
e.downcast::<ProviderError>()
|
||||
.unwrap_or_else(ProviderError::stream_decode_error)
|
||||
})?;
|
||||
if message.is_some() || usage.is_some() {
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
}
|
||||
|
||||
@@ -232,9 +232,10 @@ impl Provider for GoogleProvider {
|
||||
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))
|
||||
)?;
|
||||
let (message, usage) = message.map_err(|e| {
|
||||
e.downcast::<ProviderError>()
|
||||
.unwrap_or_else(ProviderError::stream_decode_error)
|
||||
})?;
|
||||
if message.is_some() || usage.is_some() {
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
}
|
||||
|
||||
@@ -473,9 +473,7 @@ fn stream_ollama(response: Response, mut log: RequestLog) -> Result<MessageStrea
|
||||
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))
|
||||
)?;
|
||||
let (message, usage) = message.map_err(ProviderError::from_stream_error)?;
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
yield (message, usage);
|
||||
}
|
||||
|
||||
@@ -200,7 +200,7 @@ pub fn stream_openai_compat(
|
||||
while let Some(message) = message_stream.next().await {
|
||||
let (message, usage) = message.map_err(|e|
|
||||
e.downcast::<ProviderError>()
|
||||
.unwrap_or_else(|e| ProviderError::RequestFailed(format!("Stream decode error: {e}")))
|
||||
.unwrap_or_else(ProviderError::stream_decode_error)
|
||||
)?;
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
yield (message, usage);
|
||||
@@ -223,7 +223,8 @@ pub fn stream_responses_compat(
|
||||
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}"))
|
||||
e.downcast::<ProviderError>()
|
||||
.unwrap_or_else(ProviderError::stream_decode_error)
|
||||
)?;
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
yield (message, usage);
|
||||
|
||||
Reference in New Issue
Block a user