1use 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#[derive(Clone, Debug)]
25pub struct A2aClient {
26 base_url: String,
27 http: Client,
28 request_id: Arc<AtomicU64>,
29}
30
31impl A2aClient {
32 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 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 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(¶ms)?)).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 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(¶ms)?), 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 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 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 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 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 let success = value.get("success").and_then(|v| v.as_bool()).unwrap_or(false);
189 Ok(success)
190 }
191
192 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 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
246fn find_double_newline(buf: &[u8]) -> Option<usize> {
248 buf.windows(2).position(|w| w == b"\n\n")
249}
250
251fn parse_sse_event(bytes: &[u8]) -> A2aResult<Option<StreamingEvent>> {
253 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 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}