use std::future::Future;
use std::pin::Pin;
use std::time::{Duration, Instant};
use crate::client::{A2AClient, A2ASseStream};
use crate::protocol::{A2AMessage, A2ATask, A2ATaskDetails, A2ATaskResult, AgentCard, TaskStatus};
use crate::A2AError;
type BoxedOp<'a, T> = Pin<Box<dyn Future<Output = Result<T, A2AError>> + Send + 'a>>;
#[derive(Debug, Clone)]
pub struct ResilienceConfig {
pub max_transport_retries: usize,
pub retry_base_delay: Duration,
pub max_reconnect_attempts: usize,
pub task_timeout: Duration,
}
impl Default for ResilienceConfig {
fn default() -> Self {
Self {
max_transport_retries: 2,
retry_base_delay: Duration::from_millis(50),
max_reconnect_attempts: 2,
task_timeout: Duration::from_secs(30),
}
}
}
pub struct ResilientA2AClient {
primary: A2AClient,
fallbacks: Vec<A2AClient>,
config: ResilienceConfig,
}
impl ResilientA2AClient {
pub fn new(primary: A2AClient, config: ResilienceConfig) -> Self {
Self {
primary,
fallbacks: Vec::new(),
config,
}
}
pub fn with_fallback(mut self, fallback: A2AClient) -> Self {
self.fallbacks.push(fallback);
self
}
pub fn config(&self) -> &ResilienceConfig {
&self.config
}
pub async fn get_agent_card(&self) -> Result<AgentCard, A2AError> {
self.with_fallbacks(&|c| Box::pin(async move { c.get_agent_card().await }))
.await
}
pub async fn send_task(&self, message: A2AMessage) -> Result<A2ATask, A2AError> {
self.with_fallbacks(&|c| {
let msg = message.clone();
Box::pin(async move { c.send_task(msg).await })
})
.await
}
pub async fn send_task_with_message_id(
&self,
message: A2AMessage,
message_id: &str,
) -> Result<A2ATask, A2AError> {
self.with_fallbacks(&|c| {
let msg = message.clone();
let mid = message_id.to_string();
Box::pin(async move { c.send_task_with_message_id(msg, &mid).await })
})
.await
}
pub async fn resume_task(
&self,
task_id: &str,
message: A2AMessage,
) -> Result<A2ATask, A2AError> {
self.with_fallbacks(&|c| {
let tid = task_id.to_string();
let msg = message.clone();
Box::pin(async move { c.resume_task(&tid, msg).await })
})
.await
}
pub async fn get_task(&self, task_id: &str) -> Result<A2ATask, A2AError> {
self.with_fallbacks(&|c| {
let tid = task_id.to_string();
Box::pin(async move { c.get_task(&tid).await })
})
.await
}
pub async fn get_task_details(&self, task_id: &str) -> Result<A2ATaskDetails, A2AError> {
self.with_fallbacks(&|c| {
let tid = task_id.to_string();
Box::pin(async move { c.get_task_details(&tid).await })
})
.await
}
pub async fn cancel_task(&self, task_id: &str) -> Result<A2ATask, A2AError> {
self.with_fallbacks(&|c| {
let tid = task_id.to_string();
Box::pin(async move { c.cancel_task(&tid).await })
})
.await
}
pub async fn send_task_and_wait(
&self,
message: A2AMessage,
timeout: Duration,
) -> Result<A2ATaskResult, A2AError> {
let task = self.send_task(message).await?;
self.wait_for_task(&task.id, timeout).await
}
pub async fn send_task_and_wait_with_message_id(
&self,
message: A2AMessage,
message_id: &str,
timeout: Duration,
) -> Result<A2ATaskResult, A2AError> {
let task = self.send_task_with_message_id(message, message_id).await?;
self.wait_for_task(&task.id, timeout).await
}
pub async fn wait_for_task(
&self,
task_id: &str,
timeout: Duration,
) -> Result<A2ATaskResult, A2AError> {
let start = Instant::now();
let poll_interval = Duration::from_secs(1);
loop {
let details = retry_on(
&self.primary,
self.config.max_transport_retries,
self.config.retry_base_delay,
&|c| {
let tid = task_id.to_string();
Box::pin(async move { c.get_task_details(&tid).await })
},
)
.await?;
match details.task.status {
TaskStatus::Completed => {
return details.result.ok_or_else(|| {
A2AError::Parse(format!("Task {} completed without a result", task_id))
})
}
TaskStatus::Failed => {
return Err(A2AError::Api {
code: -32000,
message: details.error.unwrap_or_else(|| "Task failed".to_string()),
})
}
TaskStatus::Cancelled => {
return Err(A2AError::Api {
code: -32000,
message: format!("Task {} was cancelled", task_id),
})
}
TaskStatus::Rejected => {
return Err(A2AError::Api {
code: -32000,
message: format!("Task {} was rejected", task_id),
})
}
TaskStatus::Expired => {
return Err(A2AError::Api {
code: -32000,
message: format!("Task {} expired", task_id),
})
}
TaskStatus::AuthRequired => {
return Err(A2AError::Api {
code: 401,
message: format!("Task {} requires authentication", task_id),
})
}
TaskStatus::InputRequired => {
return Err(A2AError::InputRequired {
task_id: task_id.to_string(),
prompt: details
.error
.unwrap_or_else(|| "Input required".to_string()),
});
}
TaskStatus::Submitted | TaskStatus::Working => {
if start.elapsed() > timeout {
return Err(A2AError::Timeout(format!(
"Task {} did not complete within {:?}; \
recover by calling wait_for_task again",
task_id, timeout
)));
}
tokio::time::sleep(poll_interval).await;
}
}
}
}
pub async fn connect_sse(&self, sse_url: &str) -> Result<A2ASseStream, A2AError> {
retry_on(
&self.primary,
self.config.max_reconnect_attempts,
self.config.retry_base_delay,
&|c| {
let url = sse_url.to_string();
Box::pin(async move { c.connect_sse(&url).await })
},
)
.await
}
pub async fn send_task_streaming(
&self,
sse_url: &str,
message: A2AMessage,
) -> Result<A2ASseStream, A2AError> {
let stream = self.connect_sse(sse_url).await?;
self.send_task(message).await?;
Ok(stream)
}
async fn with_fallbacks<F, T>(&self, op: &F) -> Result<T, A2AError>
where
F: for<'a> Fn(&'a A2AClient) -> BoxedOp<'a, T>,
{
let mut last_err: Option<A2AError> = None;
for client in std::iter::once(&self.primary).chain(self.fallbacks.iter()) {
match retry_on(
client,
self.config.max_transport_retries,
self.config.retry_base_delay,
op,
)
.await
{
Ok(value) => return Ok(value),
Err(err) if should_fallback(&err) => last_err = Some(err),
Err(err) => return Err(err),
}
}
Err(last_err.unwrap_or_else(|| {
A2AError::Http("no A2A endpoints configured for resilient client".to_string())
}))
}
}
fn is_retryable(err: &A2AError) -> bool {
matches!(err, A2AError::Http(_) | A2AError::Timeout(_))
}
fn should_fallback(err: &A2AError) -> bool {
!matches!(
err,
A2AError::Parse(_) | A2AError::Signature(_) | A2AError::InputRequired { .. }
)
}
async fn retry_on<F, T>(
client: &A2AClient,
max_retries: usize,
base_delay: Duration,
op: &F,
) -> Result<T, A2AError>
where
F: for<'a> Fn(&'a A2AClient) -> BoxedOp<'a, T>,
{
for attempt in 0..=max_retries {
match op(client).await {
Ok(value) => return Ok(value),
Err(err) if attempt < max_retries && is_retryable(&err) => {
let delay = base_delay.saturating_mul(2u32.saturating_pow(attempt as u32));
tokio::time::sleep(delay).await;
}
Err(err) => return Err(err),
}
}
unreachable!("loop covers attempts 0..={max_retries}")
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
type Handler = Arc<dyn Fn(&str, &str) -> (u16, String) + Send + Sync>;
async fn read_request(stream: &mut TcpStream) -> (String, String) {
let mut buf = vec![0u8; 4096];
let mut request = Vec::new();
let mut head_end = None;
while head_end.is_none() {
let n = stream.read(&mut buf).await.unwrap_or(0);
if n == 0 {
break;
}
request.extend_from_slice(&buf[..n]);
head_end = request.windows(4).position(|w| w == b"\r\n\r\n");
}
let head_end = head_end.expect("request head terminator");
let head = String::from_utf8_lossy(&request[..head_end]).to_string();
let body_len = head
.lines()
.find_map(|l| l.strip_prefix("Content-Length:"))
.and_then(|v| v.trim().parse::<usize>().ok())
.unwrap_or(0);
let mut body = request[head_end + 4..].to_vec();
while body.len() < body_len {
let n = stream.read(&mut buf).await.unwrap_or(0);
if n == 0 {
break;
}
body.extend_from_slice(&buf[..n]);
}
let path = head.split_whitespace().nth(1).unwrap_or("/").to_string();
(path, String::from_utf8_lossy(&body).to_string())
}
async fn write_response(stream: &mut TcpStream, status: u16, body: &str) {
let reason = match status {
200 => "OK",
500 => "Internal Server Error",
503 => "Service Unavailable",
_ => "Status",
};
let head = format!(
"HTTP/1.1 {status} {reason}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
let _ = stream.write_all(head.as_bytes()).await;
let _ = stream.write_all(body.as_bytes()).await;
let _ = stream.shutdown().await;
}
async fn spawn_server(handler: Handler) -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
loop {
let (stream, _) = match listener.accept().await {
Ok(s) => s,
Err(_) => break,
};
let handler = handler.clone();
tokio::spawn(async move {
let mut stream = stream;
let (path, body) = read_request(&mut stream).await;
let (status, response) = handler(&path, &body);
write_response(&mut stream, status, &response).await;
});
}
});
format!("http://127.0.0.1:{port}")
}
fn completed_task_response(task_id: &str, output: &str) -> String {
format!(
r#"{{"jsonrpc":"2.0","id":1,"result":{{"task":{{"id":"{task_id}","message":{{"role":"user","content":"hi"}},"status":"completed"}},"result":{{"output":"{output}"}}}}}}"#
)
}
fn working_task_response(task_id: &str) -> String {
format!(
r#"{{"jsonrpc":"2.0","id":1,"result":{{"task":{{"id":"{task_id}","message":{{"role":"user","content":"hi"}},"status":"working"}}}}}}"#
)
}
fn api_error_response() -> String {
r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32000,"message":"denied"}}"#.to_string()
}
fn fast_config() -> ResilienceConfig {
ResilienceConfig {
max_transport_retries: 2,
retry_base_delay: Duration::from_millis(10),
max_reconnect_attempts: 0,
task_timeout: Duration::from_secs(5),
}
}
#[tokio::test]
async fn transport_retry_recovers_from_transient_failures() {
let attempts = Arc::new(AtomicUsize::new(0));
let counter = attempts.clone();
let handler: Handler = Arc::new(move |_path, _body| {
let n = counter.fetch_add(1, Ordering::SeqCst);
if n < 2 {
(500, "boom".to_string())
} else {
(200, completed_task_response("task-retry", "ok"))
}
});
let base = spawn_server(handler).await;
let client = ResilientA2AClient::new(A2AClient::new(base).unwrap(), fast_config());
let task = client.send_task(A2AMessage::user("hi")).await.unwrap();
assert_eq!(task.id, "task-retry");
assert_eq!(attempts.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn api_error_is_not_retried() {
let attempts = Arc::new(AtomicUsize::new(0));
let counter = attempts.clone();
let handler: Handler = Arc::new(move |_path, _body| {
counter.fetch_add(1, Ordering::SeqCst);
(200, api_error_response())
});
let base = spawn_server(handler).await;
let client = ResilientA2AClient::new(A2AClient::new(base).unwrap(), fast_config());
let err = client.send_task(A2AMessage::user("hi")).await.unwrap_err();
assert!(matches!(err, A2AError::Api { code: -32000, .. }));
assert_eq!(attempts.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn falls_back_to_alternate_agent_when_primary_is_down() {
let primary_hits = Arc::new(AtomicUsize::new(0));
let hits = primary_hits.clone();
let primary: Handler = Arc::new(move |_path, _body| {
hits.fetch_add(1, Ordering::SeqCst);
(500, "down".to_string())
});
let primary_base = spawn_server(primary).await;
let fallback: Handler =
Arc::new(|_path, _body| (200, completed_task_response("from-fallback", "ok")));
let fallback_base = spawn_server(fallback).await;
let config = ResilienceConfig {
max_transport_retries: 1,
retry_base_delay: Duration::from_millis(5),
..fast_config()
};
let client = ResilientA2AClient::new(A2AClient::new(primary_base).unwrap(), config)
.with_fallback(A2AClient::new(fallback_base).unwrap());
let task = client.send_task(A2AMessage::user("hi")).await.unwrap();
assert_eq!(task.id, "from-fallback");
assert_eq!(primary_hits.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn does_not_fall_back_on_parse_error() {
let primary: Handler = Arc::new(|_path, _body| (200, "not-json{{{".to_string()));
let primary_base = spawn_server(primary).await;
let fallback_hits = Arc::new(AtomicUsize::new(0));
let hits = fallback_hits.clone();
let fallback: Handler = Arc::new(move |_path, _body| {
hits.fetch_add(1, Ordering::SeqCst);
(200, completed_task_response("fallback", "ok"))
});
let fallback_base = spawn_server(fallback).await;
let client = ResilientA2AClient::new(A2AClient::new(primary_base).unwrap(), fast_config())
.with_fallback(A2AClient::new(fallback_base).unwrap());
let err = client.send_task(A2AMessage::user("hi")).await.unwrap_err();
assert!(matches!(err, A2AError::Parse(_)));
assert_eq!(fallback_hits.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn wait_for_task_recovers_a_pending_task_by_id() {
let handler: Handler = Arc::new(|_path, body| {
if body.contains("tasks/get") {
(200, completed_task_response("t1", "recovered"))
} else {
(200, working_task_response("t1"))
}
});
let base = spawn_server(handler).await;
let client = ResilientA2AClient::new(A2AClient::new(base).unwrap(), fast_config());
let result = client
.send_task_and_wait(A2AMessage::user("hi"), Duration::from_secs(5))
.await
.unwrap();
assert_eq!(result.output, "recovered");
}
#[tokio::test]
async fn wait_for_task_times_out_and_exposes_task_id_for_recovery() {
let handler: Handler = Arc::new(|_path, _body| (200, working_task_response("t-stuck")));
let base = spawn_server(handler).await;
let client = ResilientA2AClient::new(A2AClient::new(base).unwrap(), fast_config());
let err = client
.wait_for_task("t-stuck", Duration::from_millis(100))
.await
.unwrap_err();
assert!(matches!(err, A2AError::Timeout(_)));
assert!(err.to_string().contains("t-stuck"));
}
#[tokio::test]
async fn connect_sse_reconnects_after_transient_failure() {
let attempts = Arc::new(AtomicUsize::new(0));
let counter = attempts.clone();
let handler: Handler = Arc::new(move |_path, _body| {
let n = counter.fetch_add(1, Ordering::SeqCst);
if n == 0 {
(500, "down".to_string())
} else {
(200, String::new())
}
});
let base = spawn_server(handler).await;
let config = ResilienceConfig {
max_reconnect_attempts: 1,
retry_base_delay: Duration::from_millis(5),
..fast_config()
};
let client = ResilientA2AClient::new(A2AClient::new(base.clone()).unwrap(), config);
let _stream = client.connect_sse(&format!("{base}/events")).await.unwrap();
assert_eq!(attempts.load(Ordering::SeqCst), 2);
}
}