1use std::collections::HashMap;
39use std::sync::Arc;
40
41use serde_json::{json, Value};
42use tokio::sync::RwLock;
43
44use lc_chains::base::BaseChain;
45
46use super::protocol::{
47 A2AErrorData, A2AMessage, A2ARequest, A2AResponse, A2ATask, A2ATaskResult, AgentCard,
48 TaskStatus,
49};
50
51#[derive(Debug, Clone)]
53struct StoredTask {
54 task: A2ATask,
55 result: Option<A2ATaskResult>,
56}
57
58const DEFAULT_MAX_TASKS: usize = 10_000;
63
64pub struct A2AServer {
71 chain: Arc<dyn BaseChain>,
73 card: AgentCard,
75 tasks: RwLock<HashMap<String, StoredTask>>,
77 max_tasks: usize,
79}
80
81impl A2AServer {
82 pub fn new(chain: Arc<dyn BaseChain>) -> Self {
84 let card = AgentCard::new(
85 chain.name(),
86 format!("Agent backed by {}", chain.name()),
87 "http://localhost:8080",
88 );
89 Self {
90 chain,
91 card,
92 tasks: RwLock::new(HashMap::new()),
93 max_tasks: DEFAULT_MAX_TASKS,
94 }
95 }
96
97 pub fn with_max_tasks(mut self, max: usize) -> Self {
99 self.max_tasks = max.max(1);
100 self
101 }
102
103 async fn evict_if_needed(&self) {
105 let mut tasks = self.tasks.write().await;
106 if tasks.len() <= self.max_tasks {
107 return;
108 }
109
110 let terminal_ids: Vec<String> = tasks
112 .iter()
113 .filter(|(_, t)| {
114 matches!(
115 t.task.status,
116 TaskStatus::Completed | TaskStatus::Failed | TaskStatus::Cancelled
117 )
118 })
119 .map(|(id, _)| id.clone())
120 .collect();
121
122 let excess = tasks.len().saturating_sub(self.max_tasks);
123 for id in terminal_ids.into_iter().take(excess) {
124 tasks.remove(&id);
125 }
126 }
127
128 pub fn with_card(mut self, card: AgentCard) -> Self {
130 self.card = card;
131 self
132 }
133
134 pub fn get_agent_card(&self) -> &AgentCard {
136 &self.card
137 }
138
139 pub async fn handle_a2a_request(&self, req: A2ARequest) -> A2AResponse {
147 match req.method.as_str() {
148 "tasks/send" => self.handle_tasks_send(req).await,
149 "tasks/get" => self.handle_tasks_get(req).await,
150 "tasks/cancel" => self.handle_tasks_cancel(req).await,
151 _ => A2AResponse::from_error_data(req.id, A2AErrorData::method_not_found()),
152 }
153 }
154
155 async fn handle_tasks_send(&self, req: A2ARequest) -> A2AResponse {
157 let params = match req.params {
158 Some(p) => p,
159 None => {
160 return A2AResponse::from_error_data(
161 req.id,
162 A2AErrorData::invalid_params("Missing params for tasks/send"),
163 )
164 }
165 };
166
167 let message: A2AMessage = match params.get("message") {
169 Some(msg_val) => serde_json::from_value(msg_val.clone()).unwrap_or_else(|_| {
170 A2AMessage::new(
171 "user",
172 msg_val
173 .get("content")
174 .and_then(|v| v.as_str())
175 .unwrap_or(""),
176 )
177 }),
178 None => {
179 A2AMessage::user(params.to_string())
181 }
182 };
183
184 let inputs: HashMap<String, Value> = {
186 let mut map = HashMap::new();
187 let input_keys = self.chain.input_keys();
189 if let Some(first_key) = input_keys.first() {
190 map.insert(
191 first_key.to_string(),
192 Value::String(message.content.clone()),
193 );
194 } else {
195 map.insert("input".to_string(), Value::String(message.content.clone()));
196 }
197 map
198 };
199
200 match self.chain.invoke(inputs).await {
202 Ok(result) => {
203 let output = result
205 .values()
206 .next()
207 .and_then(|v| v.as_str())
208 .unwrap_or("")
209 .to_string();
210
211 let task_id = uuid::Uuid::new_v4().to_string();
212 let task = A2ATask {
213 id: task_id.clone(),
214 message,
215 status: TaskStatus::Completed,
216 };
217 let task_result = A2ATaskResult::new(output);
218
219 {
221 self.tasks.write().await.insert(
222 task_id,
223 StoredTask {
224 task: task.clone(),
225 result: Some(task_result.clone()),
226 },
227 );
228 }
229 self.evict_if_needed().await;
230
231 A2AResponse::ok(
232 req.id,
233 json!({
234 "task": task,
235 "result": task_result,
236 }),
237 )
238 }
239 Err(e) => {
240 let task_id = uuid::Uuid::new_v4().to_string();
241 let task = A2ATask {
242 id: task_id.clone(),
243 message,
244 status: TaskStatus::Failed,
245 };
246
247 {
249 self.tasks.write().await.insert(
250 task_id,
251 StoredTask {
252 task: task.clone(),
253 result: None,
254 },
255 );
256 }
257 self.evict_if_needed().await;
258
259 A2AResponse::error(req.id, -32000, format!("Chain execution failed: {}", e))
260 }
261 }
262 }
263
264 async fn handle_tasks_get(&self, req: A2ARequest) -> A2AResponse {
266 let task_id = req
267 .params
268 .as_ref()
269 .and_then(|p| p.get("taskId"))
270 .and_then(|v| v.as_str())
271 .unwrap_or("");
272
273 if task_id.is_empty() {
274 return A2AResponse::from_error_data(
275 req.id,
276 A2AErrorData::invalid_params("Missing taskId parameter"),
277 );
278 }
279
280 let tasks = self.tasks.read().await;
281 match tasks.get(task_id) {
282 Some(stored) => {
283 let mut result = json!({ "task": stored.task });
284 if let Some(ref task_result) = stored.result {
285 result["result"] = json!(task_result);
286 }
287 A2AResponse::ok(req.id, result)
288 }
289 None => A2AResponse::from_error_data(
290 req.id,
291 A2AErrorData::new(-32001, format!("Task not found: {}", task_id)),
292 ),
293 }
294 }
295
296 async fn handle_tasks_cancel(&self, req: A2ARequest) -> A2AResponse {
298 let task_id = req
299 .params
300 .as_ref()
301 .and_then(|p| p.get("taskId"))
302 .and_then(|v| v.as_str())
303 .unwrap_or("");
304
305 if task_id.is_empty() {
306 return A2AResponse::from_error_data(
307 req.id,
308 A2AErrorData::invalid_params("Missing taskId parameter"),
309 );
310 }
311
312 let mut tasks = self.tasks.write().await;
313 match tasks.get_mut(task_id) {
314 Some(stored) => {
315 stored.task.status = TaskStatus::Cancelled;
316 A2AResponse::ok(req.id, json!({ "task": stored.task }))
317 }
318 None => A2AResponse::from_error_data(
319 req.id,
320 A2AErrorData::new(-32001, format!("Task not found: {}", task_id)),
321 ),
322 }
323 }
324}
325
326#[cfg(test)]
327mod tests {
328 use super::*;
329 use lc_chains::base::{BaseChain, ChainError, ChainResult};
330
331 struct EchoChain;
333
334 #[async_trait::async_trait]
335 impl BaseChain for EchoChain {
336 fn input_keys(&self) -> Vec<&str> {
337 vec!["input"]
338 }
339
340 fn output_keys(&self) -> Vec<&str> {
341 vec!["output"]
342 }
343
344 async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
345 let input = inputs.get("input").and_then(|v| v.as_str()).unwrap_or("");
346 let mut result = HashMap::new();
347 result.insert("output".to_string(), Value::String(input.to_string()));
348 Ok(result)
349 }
350
351 fn name(&self) -> &str {
352 "echo-chain"
353 }
354 }
355
356 struct FailChain;
358
359 #[async_trait::async_trait]
360 impl BaseChain for FailChain {
361 fn input_keys(&self) -> Vec<&str> {
362 vec!["input"]
363 }
364
365 fn output_keys(&self) -> Vec<&str> {
366 vec!["output"]
367 }
368
369 async fn invoke(&self, _inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
370 Err(ChainError::ExecutionError(
371 "intentional failure".to_string(),
372 ))
373 }
374
375 fn name(&self) -> &str {
376 "fail-chain"
377 }
378 }
379
380 fn echo_server() -> A2AServer {
381 A2AServer::new(Arc::new(EchoChain))
382 }
383
384 fn fail_server() -> A2AServer {
385 A2AServer::new(Arc::new(FailChain))
386 }
387
388 #[test]
389 fn get_agent_card_default() {
390 let server = echo_server();
391 let card = server.get_agent_card();
392 assert_eq!(card.name, "echo-chain");
393 assert!(card.description.contains("echo-chain"));
394 }
395
396 #[test]
397 fn get_agent_card_custom() {
398 let card = AgentCard::new("custom", "Custom agent", "http://example.com")
399 .with_capability("text-generation");
400 let server = echo_server().with_card(card);
401 let card = server.get_agent_card();
402 assert_eq!(card.name, "custom");
403 assert_eq!(card.url, "http://example.com");
404 assert_eq!(card.capabilities.len(), 1);
405 }
406
407 #[tokio::test]
408 async fn handle_tasks_send_success() {
409 let server = echo_server();
410 let msg = A2AMessage::user("hello world");
411 let req = A2ARequest::send_task(1, &msg);
412 let resp = server.handle_a2a_request(req).await;
413 assert!(!resp.is_error());
414
415 let result = resp.result.unwrap();
416 let task = result.get("task").unwrap();
417 assert_eq!(task["status"], "completed");
418
419 let task_result = result.get("result").unwrap();
420 assert_eq!(task_result["output"], "hello world");
421 }
422
423 #[tokio::test]
424 async fn handle_tasks_send_failure() {
425 let server = fail_server();
426 let msg = A2AMessage::user("hello");
427 let req = A2ARequest::send_task(2, &msg);
428 let resp = server.handle_a2a_request(req).await;
429 assert!(resp.is_error());
431
432 let err = resp.error.unwrap();
433 assert!(err.message.contains("Chain execution failed"));
434 }
435
436 #[tokio::test]
437 async fn handle_tasks_send_missing_params() {
438 let server = echo_server();
439 let req = A2ARequest::new(3, "tasks/send", None);
440 let resp = server.handle_a2a_request(req).await;
441 assert!(resp.is_error());
442 let err = resp.error.unwrap();
443 assert_eq!(err.code, -32602);
444 }
445
446 #[tokio::test]
447 async fn handle_tasks_get_missing_task_id() {
448 let server = echo_server();
449 let req = A2ARequest::new(4, "tasks/get", Some(json!({})));
450 let resp = server.handle_a2a_request(req).await;
451 assert!(resp.is_error());
452 }
453
454 #[tokio::test]
455 async fn handle_tasks_get_not_found() {
456 let server = echo_server();
457 let req = A2ARequest::get_task(5, "nonexistent-task");
458 let resp = server.handle_a2a_request(req).await;
459 assert!(resp.is_error());
460 let err = resp.error.unwrap();
461 assert!(err.message.contains("Task not found"));
462 }
463
464 #[tokio::test]
465 async fn handle_tasks_get_after_send() {
466 let server = echo_server();
467 let msg = A2AMessage::user("hello");
468 let send_req = A2ARequest::send_task(10, &msg);
469 let send_resp = server.handle_a2a_request(send_req).await;
470 let result = send_resp.result.unwrap();
471 let task_id = result["task"]["id"].as_str().unwrap().to_string();
472
473 let get_req = A2ARequest::get_task(11, &task_id);
475 let get_resp = server.handle_a2a_request(get_req).await;
476 assert!(!get_resp.is_error());
477
478 let get_result = get_resp.result.unwrap();
479 let task = get_result.get("task").unwrap();
480 assert_eq!(task["id"], task_id);
481 assert_eq!(task["status"], "completed");
482 assert!(get_result.get("result").is_some());
483 }
484
485 #[tokio::test]
486 async fn handle_tasks_cancel_nonexistent() {
487 let server = echo_server();
488 let req = A2ARequest::cancel_task(6, "task-123");
489 let resp = server.handle_a2a_request(req).await;
491 assert!(resp.is_error());
492 let err = resp.error.unwrap();
493 assert!(err.message.contains("Task not found"));
494 }
495
496 #[tokio::test]
497 async fn handle_tasks_cancel_existing_task() {
498 let server = echo_server();
499 let msg = A2AMessage::user("hello");
500 let send_req = A2ARequest::send_task(20, &msg);
501 let send_resp = server.handle_a2a_request(send_req).await;
502 let result = send_resp.result.unwrap();
503 let task_id = result["task"]["id"].as_str().unwrap().to_string();
504
505 let cancel_req = A2ARequest::cancel_task(21, &task_id);
507 let cancel_resp = server.handle_a2a_request(cancel_req).await;
508 assert!(!cancel_resp.is_error());
509
510 let cancel_result = cancel_resp.result.unwrap();
511 assert_eq!(cancel_result["task"]["status"], "cancelled");
512
513 let get_req = A2ARequest::get_task(22, &task_id);
515 let get_resp = server.handle_a2a_request(get_req).await;
516 assert!(!get_resp.is_error());
517 let get_result = get_resp.result.unwrap();
518 assert_eq!(get_result["task"]["status"], "cancelled");
519 }
520
521 #[tokio::test]
522 async fn handle_tasks_cancel_missing_task_id() {
523 let server = echo_server();
524 let req = A2ARequest::new(7, "tasks/cancel", Some(json!({})));
525 let resp = server.handle_a2a_request(req).await;
526 assert!(resp.is_error());
527 }
528
529 #[tokio::test]
530 async fn handle_unknown_method() {
531 let server = echo_server();
532 let req = A2ARequest::new(8, "foo/bar", None);
533 let resp = server.handle_a2a_request(req).await;
534 assert!(resp.is_error());
535 let err = resp.error.unwrap();
536 assert_eq!(err.code, -32601);
537 }
538
539 #[tokio::test]
540 async fn handle_tasks_send_with_raw_params() {
541 let server = echo_server();
543 let req = A2ARequest::new(9, "tasks/send", Some(json!({"query": "test query"})));
544 let resp = server.handle_a2a_request(req).await;
545 assert!(!resp.is_error());
546 }
547
548 #[tokio::test]
549 async fn handle_tasks_send_chain_with_no_input_keys() {
550 struct NoKeyChain;
552
553 #[async_trait::async_trait]
554 impl BaseChain for NoKeyChain {
555 fn input_keys(&self) -> Vec<&str> {
556 vec![]
557 }
558
559 fn output_keys(&self) -> Vec<&str> {
560 vec!["output"]
561 }
562
563 async fn invoke(
564 &self,
565 inputs: HashMap<String, Value>,
566 ) -> Result<ChainResult, ChainError> {
567 let input = inputs
568 .get("input")
569 .and_then(|v| v.as_str())
570 .unwrap_or("default");
571 let mut result = HashMap::new();
572 result.insert("output".to_string(), Value::String(input.to_string()));
573 Ok(result)
574 }
575
576 fn name(&self) -> &str {
577 "no-key-chain"
578 }
579 }
580
581 let server = A2AServer::new(Arc::new(NoKeyChain));
582 let msg = A2AMessage::user("hello");
583 let req = A2ARequest::send_task(10, &msg);
584 let resp = server.handle_a2a_request(req).await;
585 assert!(!resp.is_error());
586 }
587}