Skip to main content

vtcode_a2a/
client.rs

1//! A2A client for interacting with remote A2A agents.
2//! Provides helper methods for discovery, task operations, and streaming.
3
4use std::sync::{
5    Arc,
6    atomic::{AtomicU64, Ordering},
7};
8
9use anyhow::Context;
10use futures::{Stream, StreamExt};
11use reqwest::{Client, RequestBuilder};
12use serde_json::Value;
13use std::fmt;
14
15use crate::agent_card::AgentCard;
16use crate::errors::{A2aError, A2aErrorCode, A2aResult};
17use crate::rpc::{
18    JsonRpcRequest, ListTasksParams, METHOD_MESSAGE_SEND, METHOD_MESSAGE_STREAM, METHOD_TASKS_CANCEL, METHOD_TASKS_GET,
19    METHOD_TASKS_LIST, METHOD_TASKS_PUSH_CONFIG_GET, METHOD_TASKS_PUSH_CONFIG_SET, MessageSendParams,
20    SendStreamingMessageResponse, StreamingEvent, TaskIdParams, TaskPushNotificationConfig, TaskQueryParams,
21};
22use crate::types::Task;
23
24/// HTTP client for interacting with A2A agents
25#[derive(Clone)]
26pub struct A2aClient {
27    base_url: String,
28    http: Client,
29    request_id: Arc<AtomicU64>,
30    auth_token: Option<String>,
31}
32
33impl fmt::Debug for A2aClient {
34    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
35        formatter
36            .debug_struct("A2aClient")
37            .field("base_url", &self.base_url)
38            .field("http", &self.http)
39            .field("request_id", &self.request_id)
40            .field("auth_token", &self.auth_token.as_ref().map(|_| "REDACTED"))
41            .finish()
42    }
43}
44
45impl A2aClient {
46    /// Create a new client with default reqwest settings
47    pub fn new(base_url: impl Into<String>) -> A2aResult<Self> {
48        let http = Client::builder()
49            .redirect(reqwest::redirect::Policy::none())
50            .build()
51            .context("Failed to build HTTP client")
52            .map_err(|e| A2aError::Internal(e.to_string()))?;
53
54        Ok(Self {
55            base_url: base_url.into().trim_end_matches('/').to_string(),
56            http,
57            request_id: Arc::new(AtomicU64::new(1)),
58            auth_token: None,
59        })
60    }
61
62    /// Attach a bearer token to RPC and streaming requests.
63    pub fn with_bearer_token(mut self, token: impl Into<String>) -> Self {
64        self.auth_token = Some(token.into());
65        self
66    }
67
68    fn next_id(&self) -> String {
69        let id = self.request_id.fetch_add(1, Ordering::Relaxed);
70        format!("a2a-{id}")
71    }
72
73    fn rpc_url(&self) -> String {
74        format!("{}/a2a", self.base_url)
75    }
76
77    fn stream_url(&self) -> String {
78        format!("{}/a2a/stream", self.base_url)
79    }
80
81    fn agent_card_url(&self) -> String {
82        format!("{}/.well-known/agent-card.json", self.base_url)
83    }
84
85    fn with_authentication(&self, request: RequestBuilder) -> RequestBuilder {
86        match self.auth_token.as_deref() {
87            Some(token) => request.bearer_auth(token),
88            None => request,
89        }
90    }
91
92    /// Fetch the remote agent card
93    pub async fn agent_card(&self) -> A2aResult<AgentCard> {
94        let resp = self
95            .http
96            .get(self.agent_card_url())
97            .send()
98            .await
99            .context("Failed to fetch agent card")
100            .map_err(|e| A2aError::Internal(e.to_string()))?;
101
102        let status = resp.status();
103        if !status.is_success() {
104            return Err(A2aError::rpc(
105                A2aErrorCode::InvalidAgentResponse,
106                format!("Agent card request failed with status {status}"),
107            ));
108        }
109
110        let card = resp
111            .json::<AgentCard>()
112            .await
113            .context("Invalid agent card response")
114            .map_err(|e| A2aError::Internal(e.to_string()))?;
115        Ok(card)
116    }
117
118    /// Send a message/send RPC
119    pub async fn send_message(&self, params: MessageSendParams) -> A2aResult<Task> {
120        let result_value = self.call_rpc(METHOD_MESSAGE_SEND, Some(serde_json::to_value(&params)?)).await?;
121        let task: Task = serde_json::from_value(result_value)
122            .context("Failed to deserialize task")
123            .map_err(|e| A2aError::Internal(e.to_string()))?;
124        Ok(task)
125    }
126
127    /// Send a message/stream RPC and consume streaming events
128    pub async fn stream_message(
129        &self,
130        params: MessageSendParams,
131    ) -> A2aResult<impl Stream<Item = A2aResult<StreamingEvent>>> {
132        let req =
133            JsonRpcRequest::with_string_id(METHOD_MESSAGE_STREAM, Some(serde_json::to_value(&params)?), self.next_id());
134
135        let request = self
136            .http
137            .post(self.stream_url())
138            .header("accept", "text/event-stream")
139            .json(&req);
140        let response = self
141            .with_authentication(request)
142            .send()
143            .await
144            .context("Failed to open streaming request")
145            .map_err(|e| A2aError::Internal(e.to_string()))?;
146
147        let status = response.status();
148        if !status.is_success() {
149            return Err(A2aError::rpc(
150                A2aErrorCode::InvalidAgentResponse,
151                format!("Streaming request failed with status {status}"),
152            ));
153        }
154
155        let byte_stream = response.bytes_stream();
156
157        let stream = async_stream::try_stream! {
158            let mut buffer = Vec::new();
159            futures::pin_mut!(byte_stream);
160
161            while let Some(chunk) = byte_stream.next().await {
162                let chunk = chunk.context("Failed to read streaming chunk")
163                    .map_err(|e| A2aError::Internal(e.to_string()))?;
164                buffer.extend_from_slice(&chunk);
165
166                while let Some(pos) = find_double_newline(&buffer) {
167                    let event_bytes = buffer.drain(..pos + 2).collect::<Vec<u8>>();
168                    if let Some(event) = parse_sse_event(&event_bytes)? {
169                        yield event;
170                    }
171                }
172            }
173
174            if !buffer.is_empty() {
175                #[allow(clippy::collapsible_if, reason = "Intentional compatibility, platform, or test-only suppression.")]
176                if let Some(event) = parse_sse_event(&buffer)? {
177                    yield event;
178                }
179            }
180        };
181
182        Ok(stream)
183    }
184
185    /// Get a task by ID
186    pub async fn get_task(&self, task_id: String) -> A2aResult<Task> {
187        let params = serde_json::to_value(TaskQueryParams { id: task_id, history_length: None })?;
188        let result_value = self.call_rpc(METHOD_TASKS_GET, Some(params)).await?;
189        let task: Task = serde_json::from_value(result_value)
190            .context("Failed to deserialize task")
191            .map_err(|e| A2aError::Internal(e.to_string()))?;
192        Ok(task)
193    }
194
195    /// List tasks with filters
196    pub async fn list_tasks(&self, params: Option<ListTasksParams>) -> A2aResult<Value> {
197        let result_value = self
198            .call_rpc(METHOD_TASKS_LIST, params.map(serde_json::to_value).transpose()?)
199            .await?;
200        Ok(result_value)
201    }
202
203    /// Cancel a task
204    pub async fn cancel_task(&self, task_id: String) -> A2aResult<Task> {
205        let params = serde_json::to_value(TaskIdParams { id: task_id })?;
206        let result_value = self.call_rpc(METHOD_TASKS_CANCEL, Some(params)).await?;
207        let task: Task = serde_json::from_value(result_value)
208            .context("Failed to deserialize task")
209            .map_err(|e| A2aError::Internal(e.to_string()))?;
210        Ok(task)
211    }
212
213    /// Set push notification config
214    pub async fn set_push_config(&self, config: TaskPushNotificationConfig) -> A2aResult<bool> {
215        let value = self
216            .call_rpc(METHOD_TASKS_PUSH_CONFIG_SET, Some(serde_json::to_value(config)?))
217            .await?;
218        // Server returns {"success": true}
219        let success = value.get("success").and_then(|v| v.as_bool()).unwrap_or(false);
220        Ok(success)
221    }
222
223    /// Get push notification config
224    pub async fn get_push_config(&self, task_id: String) -> A2aResult<Option<TaskPushNotificationConfig>> {
225        let params = serde_json::to_value(TaskIdParams { id: task_id })?;
226        let value = self.call_rpc(METHOD_TASKS_PUSH_CONFIG_GET, Some(params)).await?;
227        if value.is_null() {
228            return Ok(None);
229        }
230        let cfg: TaskPushNotificationConfig = serde_json::from_value(value)
231            .context("Failed to deserialize push notification config")
232            .map_err(|e| A2aError::Internal(e.to_string()))?;
233        Ok(Some(cfg))
234    }
235
236    async fn call_rpc(&self, method: &str, params: Option<Value>) -> A2aResult<Value> {
237        let request = JsonRpcRequest::with_string_id(method, params, self.next_id());
238
239        let request = self.http.post(self.rpc_url()).json(&request);
240        let resp = self
241            .with_authentication(request)
242            .send()
243            .await
244            .context("RPC request failed")
245            .map_err(|e| A2aError::Internal(e.to_string()))?;
246
247        let status = resp.status();
248        let json: Value = resp
249            .json()
250            .await
251            .context("Failed to parse RPC response")
252            .map_err(|e| A2aError::Internal(e.to_string()))?;
253
254        if !status.is_success() {
255            return Err(A2aError::rpc(
256                A2aErrorCode::InvalidAgentResponse,
257                format!("RPC failed with status {status}: {json:?}"),
258            ));
259        }
260
261        // Deserialize JSON-RPC envelope
262        let rpc_response: crate::rpc::JsonRpcResponse = serde_json::from_value(json.clone())
263            .context("Invalid JSON-RPC response")
264            .map_err(|e| A2aError::Internal(e.to_string()))?;
265
266        if let Some(result) = rpc_response.result {
267            Ok(result)
268        } else if let Some(err) = rpc_response.error {
269            Err(A2aError::rpc(err.code.into(), err.message))
270        } else {
271            Err(A2aError::rpc(A2aErrorCode::InvalidAgentResponse, "Empty RPC response"))
272        }
273    }
274}
275
276/// Find the position of the first double newline delimiter ("\n\n")
277fn find_double_newline(buf: &[u8]) -> Option<usize> {
278    buf.windows(2).position(|w| w == b"\n\n")
279}
280
281/// Parse a single SSE event from raw bytes
282fn parse_sse_event(bytes: &[u8]) -> A2aResult<Option<StreamingEvent>> {
283    // SSE events are lines starting with "data: " and separated by blank line
284    let text = std::str::from_utf8(bytes)
285        .context("Invalid UTF-8 in SSE event")
286        .map_err(|e| A2aError::Internal(e.to_string()))?;
287
288    for line in text.lines() {
289        if let Some(payload) = line.strip_prefix("data: ") {
290            // Parse the streaming response wrapper
291            let wrapper: SendStreamingMessageResponse = serde_json::from_str(payload)
292                .context("Failed to deserialize streaming event")
293                .map_err(|e| A2aError::Internal(e.to_string()))?;
294            return Ok(Some(wrapper.event));
295        }
296    }
297
298    Ok(None)
299}
300
301#[cfg(test)]
302mod tests {
303    use super::*;
304
305    #[test]
306    fn test_find_double_newline() {
307        let data = b"data: x\n\nrest";
308        assert_eq!(find_double_newline(data), Some(7));
309    }
310
311    #[test]
312    fn test_parse_sse_event_empty() {
313        let res = parse_sse_event(b"event: ping\n\n").unwrap();
314        assert!(res.is_none());
315    }
316}