Skip to main content

adk_server/rest/controllers/
a2a.rs

1use crate::ServerConfig;
2use crate::a2a::{
3    AgentCard, Executor, ExecutorConfig, JsonRpcError, JsonRpcRequest, JsonRpcResponse, Message,
4    MessageSendParams, Task, TaskState, TaskStatus, TaskStatusUpdateEvent, TasksCancelParams,
5    TasksGetParams, UpdateEvent, build_agent_card, jsonrpc,
6};
7use adk_runner::{Runner, RunnerConfig};
8use axum::{
9    extract::State,
10    http::StatusCode,
11    response::{
12        IntoResponse, Json,
13        sse::{Event, Sse},
14    },
15};
16use futures::stream::Stream;
17use serde_json::Value;
18use std::{collections::HashMap, convert::Infallible, sync::Arc, time::Duration};
19use tokio::sync::{Mutex, Notify, RwLock, mpsc, oneshot};
20use tokio_util::sync::CancellationToken;
21
22/// In-memory task storage
23#[derive(Default)]
24pub struct TaskStore {
25    tasks: RwLock<HashMap<String, Task>>,
26}
27
28impl TaskStore {
29    pub fn new() -> Self {
30        Self::default()
31    }
32
33    pub async fn store(&self, task: Task) {
34        self.tasks.write().await.insert(task.id.clone(), task);
35    }
36
37    pub async fn get(&self, task_id: &str) -> Option<Task> {
38        self.tasks.read().await.get(task_id).cloned()
39    }
40
41    pub async fn remove(&self, task_id: &str) -> Option<Task> {
42        self.tasks.write().await.remove(task_id)
43    }
44}
45
46#[derive(Clone)]
47struct ActiveTask {
48    token: CancellationToken,
49    abort_handle: tokio::task::AbortHandle,
50    completion: Arc<Notify>,
51    context_id: String,
52}
53
54enum StreamTaskMessage {
55    Update(Box<UpdateEvent>),
56    Error(String),
57}
58
59/// Controller for A2A protocol endpoints
60#[derive(Clone)]
61pub struct A2aController {
62    config: ServerConfig,
63    agent_card: AgentCard,
64    task_store: Arc<TaskStore>,
65    active_tasks: Arc<Mutex<HashMap<String, ActiveTask>>>,
66}
67
68impl A2aController {
69    pub fn new(config: ServerConfig, base_url: &str) -> Self {
70        Self::build(config, base_url, None)
71    }
72
73    /// Create a controller whose agent card also lists the skills in `skill_index`.
74    ///
75    /// The indexed skills are appended to the agent-derived `skills[]` entries
76    /// via [`agent_skills_from_index`](crate::a2a::agent_skills_from_index) and
77    /// served at `/.well-known/agent.json`.
78    pub fn with_skill_index(
79        config: ServerConfig,
80        base_url: &str,
81        skill_index: Arc<adk_skill::SkillIndex>,
82    ) -> Self {
83        Self::build(config, base_url, Some(skill_index))
84    }
85
86    fn build(
87        config: ServerConfig,
88        base_url: &str,
89        skill_index: Option<Arc<adk_skill::SkillIndex>>,
90    ) -> Self {
91        let root_agent = config.agent_loader.root_agent();
92        let invoke_url = format!("{}/a2a", base_url.trim_end_matches('/'));
93        let mut agent_card = build_agent_card(root_agent.as_ref(), &invoke_url);
94        if let Some(skill_index) = skill_index {
95            let indexed = crate::a2a::agent_skills_from_index(&skill_index);
96            tracing::debug!(skill.count = indexed.len(), "appending indexed skills to agent card");
97            agent_card.skills.extend(indexed);
98        }
99
100        Self {
101            config,
102            agent_card,
103            task_store: Arc::new(TaskStore::new()),
104            active_tasks: Arc::new(Mutex::new(HashMap::new())),
105        }
106    }
107}
108
109fn build_runner_config(
110    controller: &A2aController,
111    root_agent: Arc<dyn adk_core::Agent>,
112    cancellation_token: Option<CancellationToken>,
113) -> Arc<RunnerConfig> {
114    let mut builder = Runner::builder()
115        .app_name(root_agent.name())
116        .agent(root_agent)
117        .session_service(controller.config.session_service.clone());
118    if let Some(ref artifact_service) = controller.config.artifact_service {
119        builder = builder.artifact_service(artifact_service.clone());
120    }
121    if let Some(ref memory_service) = controller.config.memory_service {
122        builder = builder.memory_service(memory_service.clone());
123    }
124    if let Some(ref compaction_config) = controller.config.compaction_config {
125        builder = builder.compaction_config(compaction_config.clone());
126    }
127    if let Some(ref context_cache_config) = controller.config.context_cache_config {
128        builder = builder.context_cache_config(context_cache_config.clone());
129    }
130    if let Some(ref cache_capable) = controller.config.cache_capable {
131        builder = builder.cache_capable(cache_capable.clone());
132    }
133    if let Some(cancellation_token) = cancellation_token {
134        builder = builder.cancellation_token(cancellation_token);
135    }
136    Arc::new(builder.build_config())
137}
138
139fn build_task_from_events(task_id: &str, context_id: &str, events: &[UpdateEvent]) -> Task {
140    let mut task = Task {
141        id: task_id.to_string(),
142        context_id: Some(context_id.to_string()),
143        status: TaskStatus { state: TaskState::Completed, message: None },
144        artifacts: Some(vec![]),
145        history: None,
146    };
147
148    for event in events {
149        match event {
150            UpdateEvent::TaskStatusUpdate(status) => {
151                task.status = status.status.clone();
152            }
153            UpdateEvent::TaskArtifactUpdate(artifact) => {
154                if let Some(ref mut artifacts) = task.artifacts {
155                    artifacts.push(artifact.artifact.clone());
156                }
157            }
158        }
159    }
160
161    task
162}
163
164fn build_failed_task(task_id: &str, context_id: &str, message: impl Into<String>) -> Task {
165    Task {
166        id: task_id.to_string(),
167        context_id: Some(context_id.to_string()),
168        status: TaskStatus { state: TaskState::Failed, message: Some(message.into()) },
169        artifacts: None,
170        history: None,
171    }
172}
173
174fn build_canceled_task(task_id: &str, context_id: &str) -> Task {
175    Task {
176        id: task_id.to_string(),
177        context_id: Some(context_id.to_string()),
178        status: TaskStatus { state: TaskState::Canceled, message: None },
179        artifacts: None,
180        history: None,
181    }
182}
183
184fn sanitize_internal_error(config: &ServerConfig, error: &adk_core::AdkError) -> String {
185    if config.security.expose_error_details {
186        error.to_string()
187    } else {
188        "Internal server error".to_string()
189    }
190}
191
192async fn start_task(
193    controller: &A2aController,
194    context_id: String,
195    task_id: String,
196    message: Message,
197    stream_updates: bool,
198) -> (oneshot::Receiver<adk_core::Result<Task>>, Option<mpsc::Receiver<StreamTaskMessage>>) {
199    let token = CancellationToken::new();
200    let completion = Arc::new(Notify::new());
201    let (task_tx, task_rx) = oneshot::channel();
202    let (stream_tx, stream_rx) = if stream_updates {
203        let (tx, rx) = mpsc::channel(32);
204        (Some(tx), Some(rx))
205    } else {
206        (None, None)
207    };
208
209    let root_agent = controller.config.agent_loader.root_agent();
210    let executor = Executor::new(ExecutorConfig {
211        app_name: root_agent.name().to_string(),
212        runner_config: build_runner_config(controller, root_agent, Some(token.clone())),
213        cancellation_token: Some(token.clone()),
214        #[cfg(feature = "a2a-interceptors")]
215        interceptor_chain: controller.config.interceptor_chain.clone(),
216    });
217
218    let controller_clone = controller.clone();
219    let completion_clone = completion.clone();
220    let task_id_for_task = task_id.clone();
221    let context_id_for_task = context_id.clone();
222    let stream_tx_for_task = stream_tx.clone();
223
224    let join_handle = tokio::spawn(async move {
225        let result = executor.execute(&context_id_for_task, &task_id_for_task, &message).await;
226
227        match result {
228            Ok(events) => {
229                if let Some(sender) = stream_tx_for_task {
230                    for event in &events {
231                        if sender
232                            .send(StreamTaskMessage::Update(Box::new(event.clone())))
233                            .await
234                            .is_err()
235                        {
236                            break;
237                        }
238                    }
239                }
240
241                let task = build_task_from_events(&task_id_for_task, &context_id_for_task, &events);
242                controller_clone.task_store.store(task.clone()).await;
243                let _ = task_tx.send(Ok(task));
244            }
245            Err(error) => {
246                if let Some(sender) = stream_tx_for_task {
247                    let _ = sender
248                        .send(StreamTaskMessage::Error(sanitize_internal_error(
249                            &controller_clone.config,
250                            &error,
251                        )))
252                        .await;
253                }
254                controller_clone
255                    .task_store
256                    .store(build_failed_task(
257                        &task_id_for_task,
258                        &context_id_for_task,
259                        error.to_string(),
260                    ))
261                    .await;
262                let _ = task_tx.send(Err(error));
263            }
264        }
265
266        controller_clone.active_tasks.lock().await.remove(&task_id_for_task);
267        completion_clone.notify_waiters();
268    });
269
270    controller.active_tasks.lock().await.insert(
271        task_id,
272        ActiveTask { token, abort_handle: join_handle.abort_handle(), completion, context_id },
273    );
274
275    (task_rx, stream_rx)
276}
277
278/// GET /.well-known/agent.json - Serve the agent card
279pub async fn get_agent_card(State(controller): State<A2aController>) -> impl IntoResponse {
280    Json(controller.agent_card.clone())
281}
282
283/// POST /a2a - JSON-RPC endpoint for A2A protocol
284pub async fn handle_jsonrpc(
285    State(controller): State<A2aController>,
286    Json(request): Json<JsonRpcRequest>,
287) -> impl IntoResponse {
288    if request.jsonrpc != "2.0" {
289        return Json(JsonRpcResponse::error(
290            request.id,
291            JsonRpcError::invalid_request("Invalid JSON-RPC version"),
292        ));
293    }
294
295    match request.method.as_str() {
296        jsonrpc::methods::MESSAGE_SEND => {
297            handle_message_send(&controller, request.params, request.id).await
298        }
299        jsonrpc::methods::TASKS_GET => {
300            handle_tasks_get(&controller, request.params, request.id).await
301        }
302        jsonrpc::methods::TASKS_CANCEL => {
303            handle_tasks_cancel(&controller, request.params, request.id).await
304        }
305        _ => Json(JsonRpcResponse::error(
306            request.id,
307            JsonRpcError::method_not_found(&request.method),
308        )),
309    }
310}
311
312/// POST /a2a/stream - SSE streaming endpoint for A2A protocol
313pub async fn handle_jsonrpc_stream(
314    State(controller): State<A2aController>,
315    Json(request): Json<JsonRpcRequest>,
316) -> Result<Sse<impl Stream<Item = Result<Event, Infallible>>>, (StatusCode, Json<JsonRpcResponse>)>
317{
318    if request.jsonrpc != "2.0" {
319        return Err((
320            StatusCode::BAD_REQUEST,
321            Json(JsonRpcResponse::error(
322                request.id.clone(),
323                JsonRpcError::invalid_request("Invalid JSON-RPC version"),
324            )),
325        ));
326    }
327
328    if request.method != jsonrpc::methods::MESSAGE_SEND_STREAM
329        && request.method != jsonrpc::methods::MESSAGE_SEND
330    {
331        return Err((
332            StatusCode::BAD_REQUEST,
333            Json(JsonRpcResponse::error(
334                request.id.clone(),
335                JsonRpcError::method_not_found(&request.method),
336            )),
337        ));
338    }
339
340    let params: MessageSendParams = match request.params {
341        Some(p) => serde_json::from_value(p).map_err(|e| {
342            (
343                StatusCode::BAD_REQUEST,
344                Json(JsonRpcResponse::error(
345                    request.id.clone(),
346                    JsonRpcError::invalid_params(e.to_string()),
347                )),
348            )
349        })?,
350        None => {
351            return Err((
352                StatusCode::BAD_REQUEST,
353                Json(JsonRpcResponse::error(
354                    request.id.clone(),
355                    JsonRpcError::invalid_params("Missing params"),
356                )),
357            ));
358        }
359    };
360
361    let request_id = request.id.clone();
362    let stream = create_message_stream(controller, params, request_id);
363
364    Ok(Sse::new(stream).keep_alive(
365        axum::response::sse::KeepAlive::new().interval(Duration::from_secs(15)).text("ping"),
366    ))
367}
368
369fn create_message_stream(
370    controller: A2aController,
371    params: MessageSendParams,
372    request_id: Option<Value>,
373) -> impl Stream<Item = Result<Event, Infallible>> {
374    async_stream::stream! {
375        let context_id = params
376            .message
377            .context_id
378            .clone()
379            .unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
380        let task_id = params
381            .message
382            .task_id
383            .clone()
384            .unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
385
386        let (_task_rx, maybe_stream_rx) = start_task(
387            &controller,
388            context_id.clone(),
389            task_id.clone(),
390            params.message.clone(),
391            true,
392        )
393        .await;
394
395        let Some(mut stream_rx) = maybe_stream_rx else {
396            yield Ok(Event::default().event("done").data(""));
397            return;
398        };
399
400        while let Some(message) = stream_rx.recv().await {
401            match message {
402                StreamTaskMessage::Update(event) => {
403                    let event_data = match event.as_ref() {
404                        UpdateEvent::TaskStatusUpdate(status) => {
405                            serde_json::to_string(&JsonRpcResponse::success(
406                                request_id.clone(),
407                                serde_json::to_value(status).unwrap_or_default(),
408                            ))
409                        }
410                        UpdateEvent::TaskArtifactUpdate(artifact) => {
411                            serde_json::to_string(&JsonRpcResponse::success(
412                                request_id.clone(),
413                                serde_json::to_value(artifact).unwrap_or_default(),
414                            ))
415                        }
416                    };
417
418                    if let Ok(data) = event_data {
419                        yield Ok(Event::default().data(data));
420                    }
421                }
422                StreamTaskMessage::Error(message) => {
423                    let error_response = JsonRpcResponse::error(
424                        request_id.clone(),
425                        JsonRpcError::internal_error(message),
426                    );
427                    if let Ok(data) = serde_json::to_string(&error_response) {
428                        yield Ok(Event::default().data(data));
429                    }
430                }
431            }
432        }
433
434        // Send done event
435        yield Ok(Event::default().event("done").data(""));
436    }
437}
438
439async fn handle_message_send(
440    controller: &A2aController,
441    params: Option<Value>,
442    id: Option<Value>,
443) -> Json<JsonRpcResponse> {
444    let params: MessageSendParams = match params {
445        Some(p) => match serde_json::from_value(p) {
446            Ok(p) => p,
447            Err(e) => {
448                return Json(JsonRpcResponse::error(
449                    id,
450                    JsonRpcError::invalid_params(e.to_string()),
451                ));
452            }
453        },
454        None => {
455            return Json(JsonRpcResponse::error(
456                id,
457                JsonRpcError::invalid_params("Missing params"),
458            ));
459        }
460    };
461
462    let context_id =
463        params.message.context_id.clone().unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
464    let task_id =
465        params.message.task_id.clone().unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
466
467    let (task_rx, _) =
468        start_task(controller, context_id.clone(), task_id.clone(), params.message, false).await;
469
470    match task_rx.await {
471        Ok(Ok(task)) => {
472            Json(JsonRpcResponse::success(id, serde_json::to_value(task).unwrap_or_default()))
473        }
474        Ok(Err(e)) => Json(JsonRpcResponse::error(
475            id,
476            JsonRpcError::internal_error_sanitized(
477                &e,
478                controller.config.security.expose_error_details,
479            ),
480        )),
481        Err(_) => {
482            Json(JsonRpcResponse::error(id, JsonRpcError::internal_error("Task execution aborted")))
483        }
484    }
485}
486
487async fn handle_tasks_get(
488    controller: &A2aController,
489    params: Option<Value>,
490    id: Option<Value>,
491) -> Json<JsonRpcResponse> {
492    let params: TasksGetParams = match params {
493        Some(p) => match serde_json::from_value(p) {
494            Ok(p) => p,
495            Err(e) => {
496                return Json(JsonRpcResponse::error(
497                    id,
498                    JsonRpcError::invalid_params(e.to_string()),
499                ));
500            }
501        },
502        None => {
503            return Json(JsonRpcResponse::error(
504                id,
505                JsonRpcError::invalid_params("Missing params"),
506            ));
507        }
508    };
509
510    if let Some(active_task) = controller.active_tasks.lock().await.get(&params.task_id).cloned() {
511        let task = Task {
512            id: params.task_id.clone(),
513            context_id: Some(active_task.context_id),
514            status: TaskStatus { state: TaskState::Working, message: None },
515            artifacts: None,
516            history: None,
517        };
518
519        return Json(JsonRpcResponse::success(id, serde_json::to_value(task).unwrap_or_default()));
520    }
521
522    match controller.task_store.get(&params.task_id).await {
523        Some(task) => {
524            Json(JsonRpcResponse::success(id, serde_json::to_value(task).unwrap_or_default()))
525        }
526        None => Json(JsonRpcResponse::error(
527            id,
528            JsonRpcError::internal_error(format!("Task not found: {}", params.task_id)),
529        )),
530    }
531}
532
533async fn handle_tasks_cancel(
534    controller: &A2aController,
535    params: Option<Value>,
536    id: Option<Value>,
537) -> Json<JsonRpcResponse> {
538    let params: TasksCancelParams = match params {
539        Some(p) => match serde_json::from_value(p) {
540            Ok(p) => p,
541            Err(e) => {
542                return Json(JsonRpcResponse::error(
543                    id,
544                    JsonRpcError::invalid_params(e.to_string()),
545                ));
546            }
547        },
548        None => {
549            return Json(JsonRpcResponse::error(
550                id,
551                JsonRpcError::invalid_params("Missing params"),
552            ));
553        }
554    };
555
556    let active_task = controller.active_tasks.lock().await.get(&params.task_id).cloned();
557
558    if let Some(active_task) = active_task {
559        active_task.token.cancel();
560
561        if tokio::time::timeout(Duration::from_secs(5), active_task.completion.notified())
562            .await
563            .is_err()
564        {
565            active_task.abort_handle.abort();
566            controller.active_tasks.lock().await.remove(&params.task_id);
567            controller
568                .task_store
569                .store(build_canceled_task(&params.task_id, &active_task.context_id))
570                .await;
571        }
572
573        let status = TaskStatusUpdateEvent {
574            task_id: params.task_id,
575            context_id: Some(active_task.context_id),
576            status: TaskStatus { state: TaskState::Canceled, message: None },
577            final_update: true,
578        };
579
580        return Json(JsonRpcResponse::success(
581            id,
582            serde_json::to_value(status).unwrap_or_default(),
583        ));
584    }
585
586    let status = TaskStatusUpdateEvent {
587        task_id: params.task_id,
588        context_id: Some(uuid::Uuid::new_v4().to_string()),
589        status: TaskStatus { state: TaskState::Canceled, message: None },
590        final_update: true,
591    };
592
593    Json(JsonRpcResponse::success(id, serde_json::to_value(status).unwrap_or_default()))
594}
595
596#[cfg(test)]
597mod tests {
598    use super::*;
599    use adk_core::{Agent, EventStream, InvocationContext, Result as AdkResult, SingleAgentLoader};
600    use adk_session::InMemorySessionService;
601    use async_trait::async_trait;
602    use futures::stream;
603
604    struct TestAgent;
605
606    #[async_trait]
607    impl Agent for TestAgent {
608        fn name(&self) -> &str {
609            "card_agent"
610        }
611
612        fn description(&self) -> &str {
613            "A card test agent"
614        }
615
616        fn sub_agents(&self) -> &[Arc<dyn Agent>] {
617            &[]
618        }
619
620        async fn run(&self, _ctx: Arc<dyn InvocationContext>) -> AdkResult<EventStream> {
621            Ok(Box::pin(stream::empty()))
622        }
623    }
624
625    fn test_config() -> ServerConfig {
626        let agent_loader = Arc::new(SingleAgentLoader::new(Arc::new(TestAgent)));
627        let session_service = Arc::new(InMemorySessionService::new());
628        ServerConfig::new(agent_loader, session_service)
629    }
630
631    fn skill_doc(name: &str) -> adk_skill::SkillDocument {
632        adk_skill::SkillDocument {
633            id: format!("{name}-0123456789ab"),
634            name: name.to_string(),
635            description: format!("{name} description"),
636            version: None,
637            license: None,
638            compatibility: None,
639            tags: vec!["indexed".to_string()],
640            allowed_tools: vec![],
641            references: vec![],
642            trigger: false,
643            hint: None,
644            metadata: Default::default(),
645            body: String::new(),
646            path: format!("skills/{name}.skill.md").into(),
647            hash: "0123456789ab".to_string(),
648            last_modified: None,
649            triggers: vec![],
650        }
651    }
652
653    #[test]
654    fn with_skill_index_appends_indexed_skills_to_card() {
655        let index = Arc::new(adk_skill::SkillIndex::new(vec![
656            skill_doc("skill-one"),
657            skill_doc("skill-two"),
658        ]));
659
660        let controller =
661            A2aController::with_skill_index(test_config(), "http://localhost:8080", index);
662
663        let expected = serde_json::json!([
664            {
665                "id": "card_agent",
666                "name": "card_agent",
667                "description": "A card test agent",
668                "tags": ["agent"],
669            },
670            {
671                "id": "skill-one",
672                "name": "skill-one",
673                "description": "skill-one description",
674                "tags": ["indexed"],
675            },
676            {
677                "id": "skill-two",
678                "name": "skill-two",
679                "description": "skill-two description",
680                "tags": ["indexed"],
681            },
682        ]);
683        assert_eq!(serde_json::to_value(&controller.agent_card.skills).unwrap(), expected);
684    }
685
686    #[test]
687    fn new_leaves_card_skills_agent_derived() {
688        let controller = A2aController::new(test_config(), "http://localhost:8080");
689
690        let expected = serde_json::json!([
691            {
692                "id": "card_agent",
693                "name": "card_agent",
694                "description": "A card test agent",
695                "tags": ["agent"],
696            },
697        ]);
698        assert_eq!(serde_json::to_value(&controller.agent_card.skills).unwrap(), expected);
699    }
700}