1use crate::WebhookNotifier;
10use crate::agent_card::AgentCard;
11use crate::errors::{A2aError, A2aErrorCode, A2aResult};
12use crate::rpc::{
13 JSONRPC_VERSION, JsonRpcError, JsonRpcRequest, JsonRpcResponse, ListTasksParams, METHOD_MESSAGE_SEND,
14 METHOD_MESSAGE_STREAM, METHOD_TASKS_CANCEL, METHOD_TASKS_GET, METHOD_TASKS_LIST, METHOD_TASKS_PUSH_CONFIG_GET,
15 METHOD_TASKS_PUSH_CONFIG_SET, MessageSendParams, SendStreamingMessageResponse, StreamingEvent, TaskIdParams,
16 TaskQueryParams,
17};
18use crate::task_manager::TaskManager;
19use crate::types::TaskState;
20use axum::{
21 Json, Router,
22 extract::State,
23 http::StatusCode,
24 response::{
25 IntoResponse, Response,
26 sse::{Event, Sse},
27 },
28 routing::post,
29};
30use serde_json::{Value, json};
31use std::convert::Infallible;
32use std::future::Future;
33use std::net::SocketAddr;
34use std::sync::Arc;
35use std::time::Duration;
36use tower_http::cors::CorsLayer;
37
38#[derive(Debug, Clone)]
44pub struct A2aServerState {
45 task_manager: Arc<TaskManager>,
47 agent_card: Arc<AgentCard>,
49 event_tx: Arc<tokio::sync::broadcast::Sender<StreamingEvent>>,
51 webhook_notifier: Arc<WebhookNotifier>,
53}
54
55impl A2aServerState {
56 pub fn new(task_manager: TaskManager, agent_card: AgentCard) -> Self {
58 let (event_tx, _) = tokio::sync::broadcast::channel(100);
59 Self {
60 task_manager: Arc::new(task_manager),
61 agent_card: Arc::new(agent_card),
62 event_tx: Arc::new(event_tx),
63 webhook_notifier: Arc::new(WebhookNotifier::new()),
64 }
65 }
66
67 fn vtcode_default(base_url: impl Into<String>) -> Self {
69 Self::new(TaskManager::new(), AgentCard::vtcode_default(base_url))
70 }
71}
72
73pub fn create_router(state: A2aServerState) -> Router {
79 Router::new()
80 .route("/.well-known/agent-card.json", axum::routing::get(get_agent_card))
81 .route("/a2a", post(handle_rpc))
82 .route("/a2a/stream", post(handle_stream))
83 .with_state(state)
84 .layer(CorsLayer::permissive())
85}
86
87async fn get_agent_card(State(state): State<A2aServerState>) -> Json<AgentCard> {
93 Json(state.agent_card.as_ref().clone())
94}
95
96async fn handle_rpc(
98 State(state): State<A2aServerState>,
99 Json(request): Json<JsonRpcRequest>,
100) -> Result<Json<JsonRpcResponse>, A2aErrorResponse> {
101 if request.jsonrpc != JSONRPC_VERSION {
103 return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id));
104 }
105
106 let result = match request.method.as_str() {
108 METHOD_MESSAGE_SEND => handle_message_send(&state, request.params, request.id.clone()).await,
109 METHOD_MESSAGE_STREAM => handle_message_stream(&state, request.params, request.id.clone()).await,
110 METHOD_TASKS_GET => handle_tasks_get(&state, request.params, request.id.clone()).await,
111 METHOD_TASKS_LIST => handle_tasks_list(&state, request.params, request.id.clone()).await,
112 METHOD_TASKS_CANCEL => handle_tasks_cancel(&state, request.params, request.id.clone()).await,
113 METHOD_TASKS_PUSH_CONFIG_SET => handle_push_config_set(&state, request.params, request.id.clone()).await,
114 METHOD_TASKS_PUSH_CONFIG_GET => handle_push_config_get(&state, request.params, request.id.clone()).await,
115 _ => {
116 return Err(A2aErrorResponse::method_not_found(&request.method, request.id));
117 }
118 };
119
120 match result {
121 Ok(result_value) => Ok(Json(JsonRpcResponse::success(result_value, request.id))),
122 Err(err) => Err(A2aErrorResponse::from_error(err, request.id)),
123 }
124}
125
126async fn handle_stream(State(state): State<A2aServerState>, Json(request): Json<JsonRpcRequest>) -> impl IntoResponse {
128 if request.jsonrpc != JSONRPC_VERSION {
129 return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id.clone()));
130 }
131
132 if request.method != METHOD_MESSAGE_STREAM {
133 return Err(A2aErrorResponse::method_not_found(&request.method, request.id.clone()));
134 }
135
136 let params: MessageSendParams = serde_json::from_value(request.params.unwrap_or_default())
138 .map_err(|_e| A2aErrorResponse::invalid_request("Invalid message/stream params", request.id.clone()))?;
139
140 let task_id = if let Some(task_id) = params.task_id.clone() {
142 task_id
143 } else {
144 let task = state.task_manager.create_task(params.context_id.clone()).await;
145 task.id.clone()
146 };
147
148 drop(
150 state
151 .task_manager
152 .add_message(&task_id, params.message.clone())
153 .await
154 .map_err(|e| A2aErrorResponse::from_error(e, request.id.clone()))?,
155 );
156
157 let mut rx = state.event_tx.subscribe();
159 let task_id_clone = task_id.clone();
160 let context_id = params.context_id.clone();
161 let notifier = state.webhook_notifier.clone();
162 let task_manager = state.task_manager.clone();
163
164 let stream = async_stream::stream! {
166 while let Ok(event) = rx.recv().await {
167 let matches = match &event {
169 StreamingEvent::Message { context_id: ctx, .. } => {
170 context_id.as_ref() == ctx.as_ref()
171 }
172 StreamingEvent::TaskStatus { task_id: tid, .. } => tid == &task_id_clone,
173 StreamingEvent::TaskArtifact { task_id: tid, .. } => tid == &task_id_clone,
174 _ => false,
175 };
176
177 if matches {
178 let notifier = notifier.clone();
180 let task_manager = task_manager.clone();
181 let task_id_for_hook = task_id_clone.clone();
182 let event_for_hook = event.clone();
183 drop(tokio::spawn(async move {
184 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
185 drop(notifier.send_event(&cfg, event_for_hook).await);
186 }
187 }));
188
189 let is_final = event.is_final();
190 let json = serde_json::to_string(&SendStreamingMessageResponse { event })
191 .unwrap_or_default();
192 yield Ok::<_, Infallible>(Event::default().data(json));
193
194 if is_final {
195 break;
196 }
197 }
198 }
199 };
200
201 let state_clone = state.clone();
203 let task_id_clone = task_id.clone();
204 drop(tokio::spawn(async move {
205 tokio::time::sleep(Duration::from_millis(100)).await;
207
208 drop(
210 state_clone
211 .task_manager
212 .update_status(&task_id_clone, TaskState::Working, None)
213 .await,
214 );
215
216 let status_event = StreamingEvent::TaskStatus {
218 task_id: task_id_clone.clone(),
219 context_id: params.context_id.clone(),
220 status: crate::types::TaskStatus::new(TaskState::Working),
221 kind: "status-update".to_string(),
222 r#final: false,
223 };
224 drop(state_clone.event_tx.send(status_event.clone()));
225
226 let notifier = state_clone.webhook_notifier.clone();
228 let task_manager = state_clone.task_manager.clone();
229 let task_id_for_hook = task_id_clone.clone();
230 drop(tokio::spawn(async move {
231 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
232 drop(notifier.send_event(&cfg, status_event).await);
233 }
234 }));
235
236 tokio::time::sleep(Duration::from_millis(200)).await;
238 let response_msg = crate::types::Message::agent_text("Processing your request...");
239 let message_event = StreamingEvent::Message {
240 message: response_msg,
241 context_id: params.context_id.clone(),
242 kind: "streaming-response".to_string(),
243 r#final: false,
244 };
245 drop(state_clone.event_tx.send(message_event.clone()));
246
247 let notifier = state_clone.webhook_notifier.clone();
249 let task_manager = state_clone.task_manager.clone();
250 let task_id_for_hook = task_id_clone.clone();
251 drop(tokio::spawn(async move {
252 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
253 drop(notifier.send_event(&cfg, message_event).await);
254 }
255 }));
256
257 tokio::time::sleep(Duration::from_millis(300)).await;
259 drop(
260 state_clone
261 .task_manager
262 .update_status(&task_id_clone, TaskState::Completed, None)
263 .await,
264 );
265
266 let final_status_event = StreamingEvent::TaskStatus {
268 task_id: task_id_clone,
269 context_id: params.context_id,
270 status: crate::types::TaskStatus::new(TaskState::Completed),
271 kind: "status-update".to_string(),
272 r#final: true,
273 };
274 drop(state_clone.event_tx.send(final_status_event.clone()));
275
276 let notifier = state_clone.webhook_notifier.clone();
278 let task_manager = state_clone.task_manager.clone();
279 let task_id_for_hook = final_status_event.task_id().unwrap_or_default().to_string();
280 drop(tokio::spawn(async move {
281 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
282 drop(notifier.send_event(&cfg, final_status_event).await);
283 }
284 }));
285 }));
286
287 Ok(Sse::new(Box::pin(stream)).keep_alive(
288 axum::response::sse::KeepAlive::new()
289 .interval(Duration::from_secs(15))
290 .text("keep-alive"),
291 ))
292}
293
294async fn handle_message_send(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
300 let params: MessageSendParams = serde_json::from_value(params.unwrap_or_default())
301 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid message/send params"))?;
302
303 let task_id = if let Some(task_id) = params.task_id {
305 task_id
306 } else {
307 let task = state.task_manager.create_task(params.context_id).await;
308 task.id.clone()
309 };
310
311 drop(state.task_manager.add_message(&task_id, params.message).await?);
313
314 let task = state.task_manager.update_status(&task_id, TaskState::Working, None).await?;
316
317 Ok(serde_json::to_value(task)?)
319}
320
321async fn handle_push_config_set(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
323 let config: crate::rpc::TaskPushNotificationConfig = serde_json::from_value(params.unwrap_or_default())
324 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/set params"))?;
325
326 state.task_manager.set_webhook_config(config).await?;
327
328 Ok(json!({ "success": true }))
329}
330
331async fn handle_push_config_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
333 let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
334 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/get params"))?;
335
336 let config = state.task_manager.get_webhook_config(¶ms.id).await;
337
338 Ok(serde_json::to_value(config)?)
339}
340
341fn handle_message_stream<'a>(
343 state: &'a A2aServerState,
344 params: Option<Value>,
345 id: Value,
346) -> impl Future<Output = A2aResult<Value>> + 'a {
347 handle_message_send(state, params, id)
349}
350
351async fn handle_tasks_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
353 let params: TaskQueryParams = serde_json::from_value(params.unwrap_or_default())
354 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/get params"))?;
355
356 let task = state.task_manager.get_task_or_error(¶ms.id).await?;
357
358 Ok(serde_json::to_value(task)?)
359}
360
361async fn handle_tasks_list(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
363 let params: ListTasksParams = serde_json::from_value(params.unwrap_or_default()).unwrap_or_default();
364
365 let result = state.task_manager.list_tasks(params).await;
366
367 Ok(serde_json::to_value(result)?)
368}
369
370async fn handle_tasks_cancel(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
372 let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
373 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/cancel params"))?;
374
375 let task = state.task_manager.cancel_task(¶ms.id).await?;
376
377 Ok(serde_json::to_value(task)?)
378}
379
380pub struct A2aErrorResponse {
386 response: JsonRpcResponse,
387 status_code: StatusCode,
388}
389
390impl A2aErrorResponse {
391 fn new(error: JsonRpcError, id: Value, status_code: StatusCode) -> Self {
393 Self {
394 response: JsonRpcResponse::error(error, id),
395 status_code,
396 }
397 }
398
399 fn invalid_request(message: &str, id: Value) -> Self {
401 Self::new(JsonRpcError::invalid_request(message), id, StatusCode::BAD_REQUEST)
402 }
403
404 fn method_not_found(method: &str, id: Value) -> Self {
406 Self::new(JsonRpcError::method_not_found(method), id, StatusCode::NOT_FOUND)
407 }
408
409 fn from_error(error: A2aError, id: Value) -> Self {
411 let code: i32 = error.code().into();
412 let message = error.to_string();
413 let status_code = match error {
414 A2aError::TaskNotFound(_) => StatusCode::NOT_FOUND,
415 A2aError::TaskNotCancelable(_) => StatusCode::UNPROCESSABLE_ENTITY,
416 A2aError::InvalidStateTransition { .. } => StatusCode::UNPROCESSABLE_ENTITY,
417 _ => StatusCode::INTERNAL_SERVER_ERROR,
418 };
419
420 Self::new(JsonRpcError::new(code, message), id, status_code)
421 }
422}
423
424impl IntoResponse for A2aErrorResponse {
425 fn into_response(self) -> Response {
426 (self.status_code, Json(self.response)).into_response()
427 }
428}
429
430pub async fn run(state: A2aServerState, addr: SocketAddr) -> anyhow::Result<()> {
436 let listener = tokio::net::TcpListener::bind(addr).await?;
437 tracing::info!("A2A server listening on {}", addr);
438 axum::serve(listener, create_router(state))
439 .with_graceful_shutdown(crate::shutdown_signal_logged("A2A"))
440 .await?;
441 Ok(())
442}
443
444#[cfg(test)]
445mod tests {
446 use super::*;
447
448 #[test]
449 fn test_server_state_creation() {
450 let state = A2aServerState::vtcode_default("http://localhost:8080");
451 assert_eq!(state.agent_card.name, "vtcode-agent");
452 }
453
454 #[test]
455 fn test_error_response_task_not_found() {
456 use serde_json::json;
457 let err_response = A2aErrorResponse::from_error(A2aError::TaskNotFound("test-id".to_string()), json!(1));
458 assert_eq!(err_response.status_code, StatusCode::NOT_FOUND);
459 }
460
461 #[test]
462 fn test_error_response_task_not_cancelable() {
463 use serde_json::json;
464 let err = A2aError::TaskNotCancelable("Cannot cancel completed task".to_string());
465 let err_response = A2aErrorResponse::from_error(err, json!(1));
466 assert_eq!(err_response.status_code, StatusCode::UNPROCESSABLE_ENTITY);
467 }
468
469 #[test]
470 fn test_error_response_invalid_request() {
471 use serde_json::json;
472 let err_response = A2aErrorResponse::invalid_request("Invalid JSON", json!(1));
473 assert_eq!(err_response.status_code, StatusCode::BAD_REQUEST);
474 }
475
476 #[tokio::test]
477 async fn test_server_state_with_broadcast() {
478 let state = A2aServerState::vtcode_default("http://localhost:8080");
479
480 let mut rx = state.event_tx.subscribe();
482
483 let test_event = StreamingEvent::Message {
485 message: super::super::types::Message::agent_text("Test"),
486 context_id: Some("test".to_string()),
487 kind: "streaming-response".to_string(),
488 r#final: false,
489 };
490
491 let _ignored = state.event_tx.send(test_event.clone()).expect("send event");
492
493 let received = rx.recv().await.expect("receive event");
495 assert!(!received.is_final());
496 }
497}