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#[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#[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 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
278pub async fn get_agent_card(State(controller): State<A2aController>) -> impl IntoResponse {
280 Json(controller.agent_card.clone())
281}
282
283pub 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
312pub 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 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(¶ms.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(¶ms.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(¶ms.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(¶ms.task_id);
567 controller
568 .task_store
569 .store(build_canceled_task(¶ms.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}