1use 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#[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 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 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 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 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(¶ms)?)).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 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(¶ms)?), 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 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 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 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 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 let success = value.get("success").and_then(|v| v.as_bool()).unwrap_or(false);
220 Ok(success)
221 }
222
223 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 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
276fn find_double_newline(buf: &[u8]) -> Option<usize> {
278 buf.windows(2).position(|w| w == b"\n\n")
279}
280
281fn parse_sse_event(bytes: &[u8]) -> A2aResult<Option<StreamingEvent>> {
283 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 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}