Skip to main content

lc_a2a/
server.rs

1//! A2A Server - handler functions for the Agent-to-Agent protocol.
2//!
3//! Provides `A2AServer` which holds an underlying agent (a `BaseChain` or
4//! `BaseTool`) and exposes handler functions that can be plugged into any
5//! HTTP framework (axum, actix, warp, etc.) rather than running its own
6//! server.
7//!
8//! # Endpoints
9//!
10//! - `GET /.well-known/agent.json` -> returns `AgentCard` (via `get_agent_card`)
11//! - `POST /` -> accepts `A2ARequest`, dispatches, returns `A2AResponse`
12//!   (via `handle_a2a_request`)
13//!
14//! # Task Persistence
15//!
16//! Tasks are stored in-memory using a `RwLock<HashMap>`. This allows
17//! `tasks/get` to retrieve previously created tasks and `tasks/cancel`
18//! to transition existing tasks to `Cancelled` status.
19//!
20//! For production use with persistence across restarts, wrap `A2AServer`
21//! with your own task store backed by a database.
22//!
23//! # Example
24//!
25//! ```ignore
26//! use lc_a2a::{A2AServer, AgentCard};
27//! use lc_chains::LLMChain;
28//! use std::sync::Arc;
29//!
30//! let chain = Arc::new(LLMChain::new(llm, "You are a helpful assistant"));
31//! let server = A2AServer::new(chain)
32//!     .with_card(AgentCard::new("my-agent", "A helpful agent", "http://localhost:8080"));
33//!
34//! // In your HTTP handler:
35//! let response = server.handle_a2a_request(request).await;
36//! ```
37
38use 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/// Stored task data including the result if completed.
52#[derive(Debug, Clone)]
53struct StoredTask {
54    task: A2ATask,
55    result: Option<A2ATaskResult>,
56}
57
58/// A2A Server - wraps an agent and provides handler functions.
59///
60/// The server does NOT start its own HTTP listener. Instead, it provides
61/// Default maximum number of tasks stored before LRU eviction.
62const DEFAULT_MAX_TASKS: usize = 10_000;
63
64/// `handle_a2a_request()` and `get_agent_card()` that you can call from
65/// any HTTP framework's route handler.
66///
67/// Tasks are stored in-memory so that `tasks/get` can retrieve them and
68/// `tasks/cancel` can transition their status. When the task store exceeds
69/// `max_tasks`, the oldest completed/failed/cancelled tasks are evicted first.
70pub struct A2AServer {
71    /// The underlying chain/agent.
72    chain: Arc<dyn BaseChain>,
73    /// The agent card metadata.
74    card: AgentCard,
75    /// In-memory task store.
76    tasks: RwLock<HashMap<String, StoredTask>>,
77    /// Maximum number of tasks before eviction.
78    max_tasks: usize,
79}
80
81impl A2AServer {
82    /// Create a new A2A server backed by a `BaseChain`.
83    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    /// Set the maximum number of tasks before LRU eviction.
98    pub fn with_max_tasks(mut self, max: usize) -> Self {
99        self.max_tasks = max.max(1);
100        self
101    }
102
103    /// Evict oldest completed/failed/cancelled tasks if over capacity.
104    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        // Evict terminal tasks (completed/failed/cancelled) to make room
111        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    /// Set a custom agent card.
129    pub fn with_card(mut self, card: AgentCard) -> Self {
130        self.card = card;
131        self
132    }
133
134    /// Get the agent card (for `GET /.well-known/agent.json`).
135    pub fn get_agent_card(&self) -> &AgentCard {
136        &self.card
137    }
138
139    /// Handle an incoming A2A request (for `POST /`).
140    ///
141    /// Dispatches based on the request method:
142    /// - `tasks/send` -> invoke the chain and return a task result
143    /// - `tasks/get` -> return a stored task
144    /// - `tasks/cancel` -> cancel a stored task
145    /// - unknown method -> method_not_found error
146    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    /// Handle `tasks/send`: invoke the chain with the message content.
156    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        // Extract the message from params.
168        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                // Fallback: treat the entire params as the input content.
180                A2AMessage::user(params.to_string())
181            }
182        };
183
184        // Build chain input from the message content.
185        let inputs: HashMap<String, Value> = {
186            let mut map = HashMap::new();
187            // Use the first input key expected by the chain.
188            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        // Invoke the chain.
201        match self.chain.invoke(inputs).await {
202            Ok(result) => {
203                // Extract the output text from the chain result.
204                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                // Store the task in-memory.
220                {
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                // Store the failed task in-memory.
248                {
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    /// Handle `tasks/get`: return a task by ID.
265    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    /// Handle `tasks/cancel`: cancel a task by ID.
297    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    /// A simple mock chain that echoes the input.
332    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    /// A chain that always fails.
357    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        // Chain failure now returns an error response
430        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        // Now retrieve the task via tasks/get.
474        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        // Cancelling a non-existent task now returns an error
490        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        // Cancel the task.
506        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        // Verify the task is cancelled when retrieved.
514        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        // When params has no "message" key, the entire params become the content.
542        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        /// A chain with no input keys.
551        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}