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
82 changes: 76 additions & 6 deletions src/crates/adapters/ai-adapters/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ use std::time::Duration;
use tokio::sync::mpsc;

const SEND_MESSAGE_STREAM_ATTEMPTS: usize = 10;
const TEST_CONNECTION_STREAM_ATTEMPTS: usize = 5;
const SEND_MESSAGE_RETRY_BASE_DELAY_MS: u64 = 500;

/// Streamed response result with the parsed stream and optional raw SSE receiver.
Expand Down Expand Up @@ -208,12 +209,31 @@ impl AIClient {
extra_body: Option<serde_json::Value>,
trace: Option<ModelExchangeTraceConfig>,
) -> Result<GeminiResponse> {
for attempt in 0..SEND_MESSAGE_STREAM_ATTEMPTS {
self.send_message_with_extra_body_trace_and_max_attempts(
messages,
tools,
extra_body,
trace,
SEND_MESSAGE_STREAM_ATTEMPTS,
)
.await
}

async fn send_message_with_extra_body_trace_and_max_attempts(
&self,
messages: Vec<Message>,
tools: Option<Vec<ToolDefinition>>,
extra_body: Option<serde_json::Value>,
trace: Option<ModelExchangeTraceConfig>,
max_attempts: usize,
) -> Result<GeminiResponse> {
for attempt in 0..max_attempts {
let stream_response = self
.send_message_stream_with_extra_body(
.send_message_stream_with_extra_body_and_max_attempts(
messages.clone(),
tools.clone(),
extra_body.clone(),
max_attempts,
trace.clone(),
)
.await?;
Expand All @@ -226,7 +246,7 @@ impl AIClient {
return Ok(response);
}
Err(error)
if attempt < SEND_MESSAGE_STREAM_ATTEMPTS - 1
if attempt < max_attempts - 1
&& is_transient_stream_error(&error.to_string()) =>
{
fail_aggregated_trace(
Expand All @@ -239,7 +259,7 @@ impl AIClient {
warn!(
"Retrying aggregated AI stream after transient error: attempt={}/{}, delay_ms={}, error={}",
attempt + 1,
SEND_MESSAGE_STREAM_ATTEMPTS,
max_attempts,
delay_ms,
error
);
Expand All @@ -260,12 +280,62 @@ impl AIClient {
unreachable!("send_message retry loop always returns")
}

async fn send_message_stream_with_extra_body_and_max_attempts(
&self,
messages: Vec<Message>,
tools: Option<Vec<ToolDefinition>>,
extra_body: Option<serde_json::Value>,
max_tries: usize,
trace: Option<ModelExchangeTraceConfig>,
) -> Result<StreamResponse> {
match ApiFormat::parse(&self.config.format)? {
ApiFormat::OpenAIChat => {
openai::chat::send_stream(self, messages, tools, extra_body, max_tries, trace).await
}
ApiFormat::OpenAIResponses => {
openai::responses::send_stream(self, messages, tools, extra_body, max_tries, trace)
.await
}
ApiFormat::Anthropic => {
anthropic::request::send_stream(self, messages, tools, extra_body, max_tries, trace)
.await
}
ApiFormat::Gemini => {
gemini::request::send_stream(self, messages, tools, extra_body, max_tries, trace)
.await
}
ApiFormat::GeminiCodeAssist => {
gemini::code_assist::send_stream(
self, messages, tools, extra_body, max_tries, trace,
)
.await
}
}
}

pub async fn test_connection(&self) -> Result<ConnectionTestResult> {
healthcheck::test_connection(self).await
healthcheck::test_connection(self, TEST_CONNECTION_STREAM_ATTEMPTS).await
}

pub async fn test_image_input_connection(&self) -> Result<ConnectionTestResult> {
healthcheck::test_image_input_connection(self).await
healthcheck::test_image_input_connection(self, TEST_CONNECTION_STREAM_ATTEMPTS).await
}

pub(crate) async fn send_test_message(
&self,
messages: Vec<Message>,
tools: Option<Vec<ToolDefinition>>,
max_attempts: usize,
) -> Result<GeminiResponse> {
let custom_body = self.config.custom_request_body.clone();
self.send_message_with_extra_body_trace_and_max_attempts(
messages,
tools,
custom_body,
None,
max_attempts,
)
.await
}

pub async fn list_models(&self) -> Result<Vec<RemoteModelInfo>> {
Expand Down
76 changes: 70 additions & 6 deletions src/crates/adapters/ai-adapters/src/client/healthcheck.rs
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,62 @@ pub(crate) fn image_test_response_matches_expected(response: &str) -> bool {
color_letter_stream.contains(AIClient::TEST_IMAGE_EXPECTED_CODE)
}

pub(crate) async fn test_connection(client: &AIClient) -> Result<ConnectionTestResult> {
fn connection_error_message_code(error_msg: &str) -> Option<ConnectionTestMessageCode> {
let msg = error_msg.to_ascii_lowercase();

let tls_keywords = [
"certificate",
"cert",
"tls",
"ssl",
"rustls",
"native-tls",
"handshake",
"unknownissuer",
"unknown issuer",
"invalid peer certificate",
"peer certificate",
"certificate verify",
"certificate verification",
"invalid certificate",
"self signed",
"self-signed",
"webpki",
];
if tls_keywords.iter().any(|keyword| msg.contains(keyword)) {
return Some(ConnectionTestMessageCode::TlsOrCertificateIssue);
}

let proxy_keywords = ["proxy", "tunnel", "connect tunnel", "http connect"];
if proxy_keywords.iter().any(|keyword| msg.contains(keyword)) {
return Some(ConnectionTestMessageCode::ProxyIssue);
}

let network_keywords = [
"connection failed",
"error sending request",
"dns",
"network",
"connection refused",
"connection reset",
"connection closed",
"timed out",
"timeout",
"econnreset",
"econnrefused",
"etimedout",
];
if network_keywords.iter().any(|keyword| msg.contains(keyword)) {
return Some(ConnectionTestMessageCode::NetworkIssue);
}

None
}

pub(crate) async fn test_connection(
client: &AIClient,
max_attempts: usize,
) -> Result<ConnectionTestResult> {
let start_time = std::time::Instant::now();

let test_messages = vec![Message::user(
Expand All @@ -77,7 +132,10 @@ pub(crate) async fn test_connection(client: &AIClient) -> Result<ConnectionTestR
}),
}]);

match client.send_message(test_messages, tools).await {
match client
.send_test_message(test_messages, tools, max_attempts)
.await
{
Ok(response) => {
let response_time_ms = elapsed_ms_u64(start_time);
if response.tool_calls.is_some() {
Expand Down Expand Up @@ -106,14 +164,17 @@ pub(crate) async fn test_connection(client: &AIClient) -> Result<ConnectionTestR
success: false,
response_time_ms,
model_response: None,
message_code: None,
message_code: connection_error_message_code(&error_msg),
error_details: Some(error_msg),
})
}
}
}

pub(crate) async fn test_image_input_connection(client: &AIClient) -> Result<ConnectionTestResult> {
pub(crate) async fn test_image_input_connection(
client: &AIClient,
max_attempts: usize,
) -> Result<ConnectionTestResult> {
let start_time = std::time::Instant::now();
let provider = client.config.format.to_ascii_lowercase();
let prompt = "Inspect the attached image and reply with exactly one 4-letter code for quadrant colors in TL,TR,BL,BR order using letters R,G,B,Y (R=red, G=green, B=blue, Y=yellow).";
Expand Down Expand Up @@ -160,7 +221,10 @@ pub(crate) async fn test_image_input_connection(client: &AIClient) -> Result<Con
tool_image_attachments: None,
}];

match client.send_message(test_messages, None).await {
match client
.send_test_message(test_messages, None, max_attempts)
.await
{
Ok(response) => {
if image_test_response_matches_expected(&response.text) {
Ok(ConnectionTestResult {
Expand Down Expand Up @@ -193,7 +257,7 @@ pub(crate) async fn test_image_input_connection(client: &AIClient) -> Result<Con
success: false,
response_time_ms: elapsed_ms_u64(start_time),
model_response: None,
message_code: None,
message_code: connection_error_message_code(&error_msg),
error_details: Some(error_msg),
})
}
Expand Down
13 changes: 12 additions & 1 deletion src/crates/adapters/ai-adapters/src/client/sse.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ use reqwest::{
header::{HeaderMap, RETRY_AFTER},
StatusCode,
};
use std::error::Error as StdError;
use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll};
Expand Down Expand Up @@ -75,7 +76,17 @@ fn remaining_ttft_timeout(
}

fn format_transport_error(label: &str, error: &reqwest::Error) -> String {
format!("{} connection failed: {}", label, error)
let mut message = format!("{} connection failed: {}", label, error);
let mut source = error.source();
let mut index = 1;

while let Some(cause) = source {
message.push_str(&format!("; cause {}: {}", index, cause));
source = cause.source();
index += 1;
}

message
}

fn is_retryable_http_status(status: StatusCode) -> bool {
Expand Down
3 changes: 3 additions & 0 deletions src/crates/contracts/core-types/src/ai.rs
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,9 @@ impl Message {
pub enum ConnectionTestMessageCode {
ToolCallsNotDetected,
ImageInputCheckFailed,
TlsOrCertificateIssue,
ProxyIssue,
NetworkIssue,
}

#[derive(Debug, Clone, Serialize, Deserialize)]
Expand Down
Loading
Loading