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