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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
95 changes: 84 additions & 11 deletions src/crates/adapters/ai-adapters/src/client/sse.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ use crate::client::StreamResponse;
use crate::stream::UnifiedResponse;
use crate::trace::{ModelExchangeRequestAttempt, ModelExchangeTraceConfig};
use anyhow::{anyhow, Result};
use bitfun_core_types::errors::{AiProviderError, ErrorCategory};
use chrono::{DateTime, Utc};
use futures::Stream;
use log::{debug, error, warn};
Expand Down Expand Up @@ -104,6 +105,32 @@ fn is_retryable_http_status(status: StatusCode) -> bool {
status.is_server_error() || matches!(status.as_u16(), 408 | 409 | 425 | 429)
}

fn provider_error_code(body: &str) -> Option<String> {
let value: serde_json::Value = serde_json::from_str(body).ok()?;
let error = value.get("error").unwrap_or(&value);
["code", "type", "status"].iter().find_map(|field| {
error.get(field).and_then(|value| match value {
serde_json::Value::String(value) => Some(value.clone()),
serde_json::Value::Number(value) => Some(value.to_string()),
_ => None,
})
})
}

fn http_provider_error(
label: &str,
status: StatusCode,
error_text: &str,
error_kind: &str,
) -> AiProviderError {
AiProviderError::from_parts(
format!("{} {} {}: {}", label, error_kind, status, error_text),
Some(label.to_string()),
provider_error_code(error_text),
Some(status.as_u16()),
)
}

fn exponential_retry_delay_ms(attempt: usize) -> u64 {
let shift = u32::try_from(attempt)
.unwrap_or(u32::MAX)
Expand Down Expand Up @@ -228,6 +255,7 @@ where
request_url: url.to_string(),
request_body: trace.capture_request_body.then(|| request_body.clone()),
attempt_number: attempt + 1,
round_attempt: trace.round_attempt().cloned(),
})
.await
} else {
Expand All @@ -248,22 +276,24 @@ where
.text()
.await
.unwrap_or_else(|e| format!("Failed to read error response: {}", e));
let provider_error =
http_provider_error(label, status, &error_text, "client error");
if let Some(trace) = trace.as_ref() {
trace
.sink
.request_attempt_failed(
trace_handle.as_ref(),
&format!("{} client error {}: {}", label, status, error_text),
&provider_error.to_string(),
)
.await;
}
error!("{} client error {}: {}", label, status, error_text);
return Err(anyhow!("{} client error {}: {}", label, status, error_text));
error!("{}", provider_error);
return Err(anyhow!(provider_error));
}

if status.is_success() {
debug!(
"{} request connected: {}ms, status: {}, protocol: {:?}, attempt: {}/{}",
"{} request connected: {}ms, status: {}, protocol: {:?}, transport_attempt: {}/{}",
label,
connect_time,
status,
Expand All @@ -277,9 +307,23 @@ where
.text()
.await
.unwrap_or_else(|e| format!("Failed to read error response: {}", e));
let error = anyhow!("{} error {}: {}", label, status, error_text);
let provider_error = http_provider_error(label, status, &error_text, "error");
if provider_error.category == ErrorCategory::ContextOverflow {
if let Some(trace) = trace.as_ref() {
trace
.sink
.request_attempt_failed(
trace_handle.as_ref(),
&provider_error.to_string(),
)
.await;
}
error!("{}", provider_error);
return Err(anyhow!(provider_error));
}
let error = anyhow!(provider_error);
warn!(
"{} request failed: {}ms, attempt {}/{}, error: {}",
"{} request failed: {}ms, transport_attempt {}/{}, error: {}",
label,
connect_time,
attempt + 1,
Expand All @@ -303,7 +347,7 @@ where
if attempt < max_tries - 1 {
let delay_ms = retry_delay_ms(attempt, &headers, status);
debug!(
"Retrying {} after {}ms (attempt {}, status {})",
"Retrying {} after {}ms (transport_attempt {}, status {})",
label,
delay_ms,
attempt + 2,
Expand All @@ -319,7 +363,7 @@ where
let error_msg = format_transport_error(label, &e);
let error = anyhow!("{}", error_msg);
warn!(
"{} request failed: {}ms, attempt {}/{}, error: {}",
"{} request failed: {}ms, transport_attempt {}/{}, error: {}",
label,
connect_time,
attempt + 1,
Expand All @@ -337,7 +381,7 @@ where
if attempt < max_tries - 1 {
let delay_ms = exponential_retry_delay_ms(attempt);
debug!(
"Retrying {} after {}ms (attempt {})",
"Retrying {} after {}ms (transport_attempt {})",
label,
delay_ms,
attempt + 2
Expand All @@ -351,7 +395,7 @@ where
let error_msg = format_ttft_timeout_error(label, ttft_timeout);
let error = anyhow!("{}", error_msg);
warn!(
"{} request failed: {}ms, attempt {}/{}, error: {}",
"{} request failed: {}ms, transport_attempt {}/{}, error: {}",
label,
connect_time,
attempt + 1,
Expand All @@ -369,7 +413,7 @@ where
if attempt < max_tries - 1 {
let delay_ms = exponential_retry_delay_ms(attempt);
debug!(
"Retrying {} after {}ms (attempt {})",
"Retrying {} after {}ms (transport_attempt {})",
label,
delay_ms,
attempt + 2
Expand Down Expand Up @@ -419,6 +463,35 @@ mod tests {
Arc,
};

#[test]
fn http_error_uses_structured_code_before_generic_message() {
let error = http_provider_error(
"OpenAI Responses API",
StatusCode::BAD_REQUEST,
r#"{"error":{"code":"context_length_exceeded","message":"Request failed"}}"#,
"client error",
);

assert_eq!(error.category, ErrorCategory::ContextOverflow);
assert_eq!(
error.provider_code.as_deref(),
Some("context_length_exceeded")
);
assert_eq!(error.http_status, Some(400));
}

#[test]
fn no_body_bad_request_is_not_assumed_to_be_context_overflow() {
let error = http_provider_error(
"OpenAI Responses API",
StatusCode::BAD_REQUEST,
"400 status code (no body)",
"client error",
);

assert_eq!(error.category, ErrorCategory::InvalidRequest);
}

#[test]
fn format_ttft_timeout_error_includes_timeout_seconds() {
let message = format_ttft_timeout_error(
Expand Down
2 changes: 1 addition & 1 deletion src/crates/adapters/ai-adapters/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ pub use model_selector::{
pub use stream::{UnifiedResponse, UnifiedTokenUsage, UnifiedToolCall};
pub use trace::{
ModelExchangeRequestAttempt, ModelExchangeRequestTraceHandle, ModelExchangeResponseTrace,
ModelExchangeTraceConfig, ModelExchangeTraceSink,
ModelExchangeRoundAttempt, ModelExchangeTraceConfig, ModelExchangeTraceSink,
};
pub use types::{
resolve_request_url, AIConfig, ConnectionTestMessageCode, ConnectionTestResult, GeminiResponse,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ use crate::stream::types::anthropic::{
};
use crate::stream::types::unified::UnifiedResponse;
use anyhow::{anyhow, Result};
use bitfun_core_types::errors::AiProviderError;
use eventsource_stream::Eventsource;
use log::{error, trace};
use reqwest::Response;
Expand Down Expand Up @@ -100,11 +101,11 @@ pub async fn handle_anthropic_stream(
let _ = tx.send(format!("[{}] {}", event_type, data));
}

if let Some(error_msg) = format_provider_error_from_sse_message(&event_type, &data) {
if let Some(provider_error) = provider_error_from_sse_message(&event_type, &data) {
stats.increment("error:provider_message");
stats.log_summary("provider_error_message_received");
error!("{}", error_msg);
let _ = tx_event.send(Err(anyhow!(error_msg)));
error!("{}", provider_error);
let _ = tx_event.send(Err(anyhow!(provider_error)));
return;
}

Expand Down Expand Up @@ -223,7 +224,14 @@ pub async fn handle_anthropic_stream(
};
stats.increment("error:api");
stats.log_summary("error_event_received");
let _ = tx_event.send(Err(anyhow!(String::from(sse_error.error))));
let code = sse_error.error.error_type.clone();
let provider_error = AiProviderError::from_parts(
String::from(sse_error.error),
Some("anthropic".to_string()),
Some(code),
None,
);
let _ = tx_event.send(Err(anyhow!(provider_error)));
return;
}
"message_stop" => {
Expand All @@ -240,7 +248,7 @@ pub async fn handle_anthropic_stream(
}
}

fn format_provider_error_from_sse_message(event_type: &str, data: &str) -> Option<String> {
fn provider_error_from_sse_message(event_type: &str, data: &str) -> Option<AiProviderError> {
if event_type != "message" {
return None;
}
Expand Down Expand Up @@ -269,7 +277,12 @@ fn format_provider_error_from_sse_message(event_type: &str, data: &str) -> Optio
formatted.push_str(&format!(", request_id={}", request_id));
}

Some(formatted)
Some(AiProviderError::from_parts(
formatted,
Some("anthropic_compatible".to_string()),
Some(code),
None,
))
}

fn should_trace_anthropic_sse_event(event_type: &str, _data: &str) -> bool {
Expand Down Expand Up @@ -376,28 +389,31 @@ fn emit_normalized_response(
#[cfg(test)]
mod tests {
use super::{
format_provider_error_from_sse_message, should_log_full_stream_events,
provider_error_from_sse_message, should_log_full_stream_events,
should_trace_anthropic_sse_event, should_trace_unified_response,
};
use crate::stream::types::unified::{UnifiedResponse, UnifiedToolCall};
use bitfun_core_types::errors::ErrorCategory;

#[test]
fn extracts_glm_business_error_from_message_event() {
let raw = r#"{"error":{"code":"1113","message":"余额不足或无可用资源包,请充值。"},"request_id":"20260425142416"}"#;

let formatted = format_provider_error_from_sse_message("message", raw).unwrap();
let error = provider_error_from_sse_message("message", raw).unwrap();
let formatted = error.message;

assert!(formatted.contains("Provider error"));
assert!(formatted.contains("code=1113"));
assert!(formatted.contains("余额不足或无可用资源包"));
assert!(formatted.contains("request_id=20260425142416"));
assert_eq!(error.category, ErrorCategory::ProviderQuota);
}

#[test]
fn ignores_regular_anthropic_delta_events() {
let raw = r#"{"type":"message_delta","delta":{"stop_reason":null}}"#;

assert!(format_provider_error_from_sse_message("message_delta", raw).is_none());
assert!(provider_error_from_sse_message("message_delta", raw).is_none());
}

#[test]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ use super::{next_stream_item, StreamTimeoutController, StreamTimeoutStage, Timed
use crate::stream::types::gemini::GeminiSSEData;
use crate::stream::types::unified::UnifiedResponse;
use anyhow::{anyhow, Result};
use bitfun_core_types::errors::AiProviderError;
use eventsource_stream::Eventsource;
use log::{error, trace};
use reqwest::Response;
Expand Down Expand Up @@ -76,6 +77,24 @@ fn extract_api_error_message(event_json: &Value) -> Option<String> {
Some("Gemini streaming request failed".to_string())
}

fn extract_api_error(event_json: &Value) -> Option<AiProviderError> {
let error = event_json.get("error")?;
let code = error
.get("status")
.or_else(|| error.get("code"))
.and_then(|value| match value {
Value::String(value) => Some(value.clone()),
Value::Number(value) => Some(value.to_string()),
_ => None,
});
Some(AiProviderError::from_parts(
extract_api_error_message(event_json)?,
Some("gemini".to_string()),
code,
None,
))
}

pub async fn handle_gemini_stream(
response: Response,
tx_event: mpsc::UnboundedSender<Result<UnifiedResponse>>,
Expand Down Expand Up @@ -157,12 +176,15 @@ pub async fn handle_gemini_stream(
}
};

if let Some(message) = extract_api_error_message(&event_json) {
let error_msg = format!("Gemini SSE API error: {}, data: {}", message, raw);
if let Some(mut provider_error) = extract_api_error(&event_json) {
provider_error.message = format!(
"Gemini SSE API error: {}, data: {}",
provider_error.message, raw
);
stats.increment("error:api");
stats.log_summary("sse_api_error");
error!("{}", error_msg);
let _ = tx_event.send(Err(anyhow!(error_msg)));
error!("{}", provider_error);
let _ = tx_event.send(Err(anyhow!(provider_error)));
return;
}

Expand Down Expand Up @@ -211,8 +233,9 @@ pub async fn handle_gemini_stream(

#[cfg(test)]
mod tests {
use super::GeminiToolCallState;
use super::{extract_api_error, GeminiToolCallState};
use crate::stream::types::unified::UnifiedToolCall;
use bitfun_core_types::errors::ErrorCategory;

#[test]
fn reuses_active_tool_id_by_omitting_follow_up_ids() {
Expand Down Expand Up @@ -325,4 +348,19 @@ mod tests {

assert_ne!(first.id, second.id);
}

#[test]
fn classifies_context_overflow_from_gemini_error_message() {
let event = serde_json::json!({
"error": {
"code": 400,
"status": "INVALID_ARGUMENT",
"message": "The input token count exceeds the maximum number of tokens allowed"
}
});

let error = extract_api_error(&event).expect("provider error");
assert_eq!(error.category, ErrorCategory::ContextOverflow);
assert_eq!(error.provider_code.as_deref(), Some("INVALID_ARGUMENT"));
}
}
Loading