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::{Request, State},
23 http::{
24 HeaderMap, HeaderValue, Method, StatusCode,
25 header::{AUTHORIZATION, CONTENT_TYPE, WWW_AUTHENTICATE},
26 },
27 middleware::{self, Next},
28 response::{
29 IntoResponse, Response,
30 sse::{Event, Sse},
31 },
32 routing::post,
33};
34use serde_json::{Value, json};
35use std::convert::Infallible;
36use std::fmt;
37use std::future::Future;
38use std::net::SocketAddr;
39use std::sync::Arc;
40use std::time::Duration;
41use tower_http::cors::CorsLayer;
42use uuid::Uuid;
43
44#[derive(Clone)]
46struct AuthToken(Arc<str>);
47
48impl fmt::Debug for AuthToken {
49 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
50 formatter.write_str("AuthToken(REDACTED)")
51 }
52}
53
54impl AuthToken {
55 fn generated() -> Self {
56 Self(Arc::from(Uuid::new_v4().to_string()))
57 }
58
59 fn from_value(value: String) -> anyhow::Result<Self> {
60 if value.is_empty() || !value.is_ascii() || value.chars().any(char::is_whitespace) {
61 anyhow::bail!("A2A authentication token must be non-empty ASCII text without whitespace");
62 }
63
64 Ok(Self(Arc::from(value)))
65 }
66
67 fn as_str(&self) -> &str {
68 &self.0
69 }
70
71 fn matches(&self, presented: &str) -> bool {
72 let expected = self.0.as_bytes();
73 let presented = presented.as_bytes();
74 let mut difference = expected.len() ^ presented.len();
75 let comparison_len = expected.len().max(presented.len());
76
77 for index in 0..comparison_len {
78 let expected_byte = expected.get(index).copied().unwrap_or_default();
79 let presented_byte = presented.get(index).copied().unwrap_or_default();
80 difference |= usize::from(expected_byte ^ presented_byte);
81 }
82
83 difference == 0
84 }
85}
86
87#[derive(Debug, Clone)]
93pub struct A2aServerState {
94 task_manager: Arc<TaskManager>,
96 agent_card: Arc<AgentCard>,
98 event_tx: Arc<tokio::sync::broadcast::Sender<StreamingEvent>>,
100 webhook_notifier: Arc<WebhookNotifier>,
102 auth_token: AuthToken,
104}
105
106impl A2aServerState {
107 pub fn new(task_manager: TaskManager, agent_card: AgentCard) -> Self {
109 Self::from_auth_token(task_manager, agent_card, AuthToken::generated())
110 }
111
112 pub fn new_with_auth_token(
114 task_manager: TaskManager,
115 agent_card: AgentCard,
116 auth_token: impl Into<String>,
117 ) -> anyhow::Result<Self> {
118 let auth_token = AuthToken::from_value(auth_token.into())?;
119 Ok(Self::from_auth_token(task_manager, agent_card, auth_token))
120 }
121
122 fn from_auth_token(task_manager: TaskManager, agent_card: AgentCard, auth_token: AuthToken) -> Self {
123 let (event_tx, _) = tokio::sync::broadcast::channel(100);
124 Self {
125 task_manager: Arc::new(task_manager),
126 agent_card: Arc::new(agent_card.with_bearer_auth()),
127 event_tx: Arc::new(event_tx),
128 webhook_notifier: Arc::new(WebhookNotifier::new()),
129 auth_token,
130 }
131 }
132
133 pub fn auth_token(&self) -> &str {
135 self.auth_token.as_str()
136 }
137
138 fn vtcode_default(base_url: impl Into<String>) -> Self {
140 Self::new(TaskManager::new(), AgentCard::vtcode_default(base_url))
141 }
142}
143
144pub fn create_router(state: A2aServerState) -> Router {
150 let protected_routes = Router::new()
151 .route("/a2a", post(handle_rpc))
152 .route("/a2a/stream", post(handle_stream))
153 .layer(middleware::from_fn_with_state(state.clone(), require_bearer_auth));
154
155 Router::new()
156 .route("/.well-known/agent-card.json", axum::routing::get(get_agent_card))
157 .merge(protected_routes)
158 .with_state(state)
159 .layer(
162 CorsLayer::new()
163 .allow_methods([Method::POST])
164 .allow_headers([AUTHORIZATION, CONTENT_TYPE]),
165 )
166}
167
168async fn require_bearer_auth(State(state): State<A2aServerState>, request: Request, next: Next) -> Response {
169 if has_valid_bearer_auth(&state, request.headers()) {
170 next.run(request).await
171 } else {
172 let mut response = (
173 StatusCode::UNAUTHORIZED,
174 Json(JsonRpcResponse::error(JsonRpcError::new(-32600, "Authentication required"), Value::Null)),
175 )
176 .into_response();
177 drop(
178 response
179 .headers_mut()
180 .insert(WWW_AUTHENTICATE, HeaderValue::from_static("Bearer")),
181 );
182 response
183 }
184}
185
186fn has_valid_bearer_auth(state: &A2aServerState, headers: &HeaderMap) -> bool {
187 if headers.get_all(AUTHORIZATION).iter().count() != 1 {
188 return false;
189 }
190
191 let Some(value) = headers.get(AUTHORIZATION).and_then(|value| value.to_str().ok()) else {
192 return false;
193 };
194
195 let Some((scheme, token)) = value.split_once(' ') else {
196 return false;
197 };
198
199 scheme.eq_ignore_ascii_case("Bearer")
200 && !token.is_empty()
201 && !token.chars().any(char::is_whitespace)
202 && state.auth_token.matches(token)
203}
204
205async fn get_agent_card(State(state): State<A2aServerState>) -> Json<AgentCard> {
211 Json(state.agent_card.as_ref().clone())
212}
213
214async fn handle_rpc(
216 State(state): State<A2aServerState>,
217 Json(request): Json<JsonRpcRequest>,
218) -> Result<Json<JsonRpcResponse>, A2aErrorResponse> {
219 if request.jsonrpc != JSONRPC_VERSION {
221 return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id));
222 }
223
224 let result = match request.method.as_str() {
226 METHOD_MESSAGE_SEND => handle_message_send(&state, request.params, request.id.clone()).await,
227 METHOD_MESSAGE_STREAM => handle_message_stream(&state, request.params, request.id.clone()).await,
228 METHOD_TASKS_GET => handle_tasks_get(&state, request.params, request.id.clone()).await,
229 METHOD_TASKS_LIST => handle_tasks_list(&state, request.params, request.id.clone()).await,
230 METHOD_TASKS_CANCEL => handle_tasks_cancel(&state, request.params, request.id.clone()).await,
231 METHOD_TASKS_PUSH_CONFIG_SET => handle_push_config_set(&state, request.params, request.id.clone()).await,
232 METHOD_TASKS_PUSH_CONFIG_GET => handle_push_config_get(&state, request.params, request.id.clone()).await,
233 _ => {
234 return Err(A2aErrorResponse::method_not_found(&request.method, request.id));
235 }
236 };
237
238 match result {
239 Ok(result_value) => Ok(Json(JsonRpcResponse::success(result_value, request.id))),
240 Err(err) => Err(A2aErrorResponse::from_error(err, request.id)),
241 }
242}
243
244async fn handle_stream(State(state): State<A2aServerState>, Json(request): Json<JsonRpcRequest>) -> impl IntoResponse {
246 if request.jsonrpc != JSONRPC_VERSION {
247 return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id.clone()));
248 }
249
250 if request.method != METHOD_MESSAGE_STREAM {
251 return Err(A2aErrorResponse::method_not_found(&request.method, request.id.clone()));
252 }
253
254 let params: MessageSendParams = serde_json::from_value(request.params.unwrap_or_default())
256 .map_err(|_e| A2aErrorResponse::invalid_request("Invalid message/stream params", request.id.clone()))?;
257
258 let task_id = if let Some(task_id) = params.task_id.clone() {
260 task_id
261 } else {
262 let task = state.task_manager.create_task(params.context_id.clone()).await;
263 task.id.clone()
264 };
265
266 drop(
268 state
269 .task_manager
270 .add_message(&task_id, params.message.clone())
271 .await
272 .map_err(|e| A2aErrorResponse::from_error(e, request.id.clone()))?,
273 );
274
275 let mut rx = state.event_tx.subscribe();
277 let task_id_clone = task_id.clone();
278 let context_id = params.context_id.clone();
279 let notifier = state.webhook_notifier.clone();
280 let task_manager = state.task_manager.clone();
281
282 let stream = async_stream::stream! {
284 while let Ok(event) = rx.recv().await {
285 let matches = match &event {
287 StreamingEvent::Message { context_id: ctx, .. } => {
288 context_id.as_ref() == ctx.as_ref()
289 }
290 StreamingEvent::TaskStatus { task_id: tid, .. } => tid == &task_id_clone,
291 StreamingEvent::TaskArtifact { task_id: tid, .. } => tid == &task_id_clone,
292 _ => false,
293 };
294
295 if matches {
296 let notifier = notifier.clone();
298 let task_manager = task_manager.clone();
299 let task_id_for_hook = task_id_clone.clone();
300 let event_for_hook = event.clone();
301 drop(tokio::spawn(async move {
302 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
303 drop(notifier.send_event(&cfg, event_for_hook).await);
304 }
305 }));
306
307 let is_final = event.is_final();
308 let json = serde_json::to_string(&SendStreamingMessageResponse { event })
309 .unwrap_or_default();
310 yield Ok::<_, Infallible>(Event::default().data(json));
311
312 if is_final {
313 break;
314 }
315 }
316 }
317 };
318
319 let state_clone = state.clone();
321 let task_id_clone = task_id.clone();
322 drop(tokio::spawn(async move {
323 tokio::time::sleep(Duration::from_millis(100)).await;
325
326 drop(
328 state_clone
329 .task_manager
330 .update_status(&task_id_clone, TaskState::Working, None)
331 .await,
332 );
333
334 let status_event = StreamingEvent::TaskStatus {
336 task_id: task_id_clone.clone(),
337 context_id: params.context_id.clone(),
338 status: crate::types::TaskStatus::new(TaskState::Working),
339 kind: "status-update".to_string(),
340 r#final: false,
341 };
342 drop(state_clone.event_tx.send(status_event.clone()));
343
344 let notifier = state_clone.webhook_notifier.clone();
346 let task_manager = state_clone.task_manager.clone();
347 let task_id_for_hook = task_id_clone.clone();
348 drop(tokio::spawn(async move {
349 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
350 drop(notifier.send_event(&cfg, status_event).await);
351 }
352 }));
353
354 tokio::time::sleep(Duration::from_millis(200)).await;
356 let response_msg = crate::types::Message::agent_text("Processing your request...");
357 let message_event = StreamingEvent::Message {
358 message: response_msg,
359 context_id: params.context_id.clone(),
360 kind: "streaming-response".to_string(),
361 r#final: false,
362 };
363 drop(state_clone.event_tx.send(message_event.clone()));
364
365 let notifier = state_clone.webhook_notifier.clone();
367 let task_manager = state_clone.task_manager.clone();
368 let task_id_for_hook = task_id_clone.clone();
369 drop(tokio::spawn(async move {
370 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
371 drop(notifier.send_event(&cfg, message_event).await);
372 }
373 }));
374
375 tokio::time::sleep(Duration::from_millis(300)).await;
377 drop(
378 state_clone
379 .task_manager
380 .update_status(&task_id_clone, TaskState::Completed, None)
381 .await,
382 );
383
384 let final_status_event = StreamingEvent::TaskStatus {
386 task_id: task_id_clone,
387 context_id: params.context_id,
388 status: crate::types::TaskStatus::new(TaskState::Completed),
389 kind: "status-update".to_string(),
390 r#final: true,
391 };
392 drop(state_clone.event_tx.send(final_status_event.clone()));
393
394 let notifier = state_clone.webhook_notifier.clone();
396 let task_manager = state_clone.task_manager.clone();
397 let task_id_for_hook = final_status_event.task_id().unwrap_or_default().to_string();
398 drop(tokio::spawn(async move {
399 if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
400 drop(notifier.send_event(&cfg, final_status_event).await);
401 }
402 }));
403 }));
404
405 Ok(Sse::new(Box::pin(stream)).keep_alive(
406 axum::response::sse::KeepAlive::new()
407 .interval(Duration::from_secs(15))
408 .text("keep-alive"),
409 ))
410}
411
412async fn handle_message_send(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
418 let params: MessageSendParams = serde_json::from_value(params.unwrap_or_default())
419 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid message/send params"))?;
420
421 let task_id = if let Some(task_id) = params.task_id {
423 task_id
424 } else {
425 let task = state.task_manager.create_task(params.context_id).await;
426 task.id.clone()
427 };
428
429 drop(state.task_manager.add_message(&task_id, params.message).await?);
431
432 let task = state.task_manager.update_status(&task_id, TaskState::Working, None).await?;
434
435 Ok(serde_json::to_value(task)?)
437}
438
439async fn handle_push_config_set(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
441 let config: crate::rpc::TaskPushNotificationConfig = serde_json::from_value(params.unwrap_or_default())
442 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/set params"))?;
443
444 state.task_manager.set_webhook_config(config).await?;
445
446 Ok(json!({ "success": true }))
447}
448
449async fn handle_push_config_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
451 let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
452 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/get params"))?;
453
454 let config = state.task_manager.get_webhook_config(¶ms.id).await;
455
456 Ok(serde_json::to_value(config)?)
457}
458
459fn handle_message_stream<'a>(
461 state: &'a A2aServerState,
462 params: Option<Value>,
463 id: Value,
464) -> impl Future<Output = A2aResult<Value>> + 'a {
465 handle_message_send(state, params, id)
467}
468
469async fn handle_tasks_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
471 let params: TaskQueryParams = serde_json::from_value(params.unwrap_or_default())
472 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/get params"))?;
473
474 let task = state
475 .task_manager
476 .get_task_or_error_with_history(¶ms.id, params.history_length.unwrap_or_default() as usize)
477 .await?;
478
479 Ok(serde_json::to_value(task)?)
480}
481
482async fn handle_tasks_list(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
484 let params = match params {
485 Some(params) => serde_json::from_value(params)
486 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/list params"))?,
487 None => ListTasksParams::default(),
488 };
489
490 let result = state.task_manager.list_tasks(params).await;
491
492 Ok(serde_json::to_value(result)?)
493}
494
495async fn handle_tasks_cancel(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
497 let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
498 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/cancel params"))?;
499
500 let task = state.task_manager.cancel_task(¶ms.id).await?;
501
502 Ok(serde_json::to_value(task)?)
503}
504
505pub struct A2aErrorResponse {
511 response: JsonRpcResponse,
512 status_code: StatusCode,
513}
514
515impl A2aErrorResponse {
516 fn new(error: JsonRpcError, id: Value, status_code: StatusCode) -> Self {
518 Self {
519 response: JsonRpcResponse::error(error, id),
520 status_code,
521 }
522 }
523
524 fn invalid_request(message: &str, id: Value) -> Self {
526 Self::new(JsonRpcError::invalid_request(message), id, StatusCode::BAD_REQUEST)
527 }
528
529 fn method_not_found(method: &str, id: Value) -> Self {
531 Self::new(JsonRpcError::method_not_found(method), id, StatusCode::NOT_FOUND)
532 }
533
534 fn from_error(error: A2aError, id: Value) -> Self {
536 let code: i32 = error.code().into();
537 let message = error.to_string();
538 let status_code = match &error {
539 A2aError::RpcError { code, .. } => match code {
540 A2aErrorCode::InvalidRequest | A2aErrorCode::InvalidParams | A2aErrorCode::JsonParseError => {
541 StatusCode::BAD_REQUEST
542 }
543 A2aErrorCode::MethodNotFound => StatusCode::NOT_FOUND,
544 _ => StatusCode::INTERNAL_SERVER_ERROR,
545 },
546 A2aError::TaskNotFound(_) => StatusCode::NOT_FOUND,
547 A2aError::TaskNotCancelable(_) => StatusCode::UNPROCESSABLE_ENTITY,
548 A2aError::InvalidStateTransition { .. } => StatusCode::UNPROCESSABLE_ENTITY,
549 _ => StatusCode::INTERNAL_SERVER_ERROR,
550 };
551
552 Self::new(JsonRpcError::new(code, message), id, status_code)
553 }
554}
555
556impl IntoResponse for A2aErrorResponse {
557 fn into_response(self) -> Response {
558 (self.status_code, Json(self.response)).into_response()
559 }
560}
561
562pub async fn run(state: A2aServerState, addr: SocketAddr) -> anyhow::Result<()> {
568 let listener = tokio::net::TcpListener::bind(addr).await?;
569 tracing::info!("A2A server listening on {}", addr);
570 axum::serve(listener, create_router(state))
571 .with_graceful_shutdown(crate::shutdown_signal_logged("A2A"))
572 .await?;
573 Ok(())
574}
575
576#[cfg(test)]
577mod tests {
578 use super::*;
579
580 #[test]
581 fn test_server_state_creation() {
582 let state = A2aServerState::vtcode_default("http://localhost:8080");
583 assert_eq!(state.agent_card.name, "vtcode-agent");
584 }
585
586 #[test]
587 fn test_error_response_task_not_found() {
588 use serde_json::json;
589 let err_response = A2aErrorResponse::from_error(A2aError::TaskNotFound("test-id".to_string()), json!(1));
590 assert_eq!(err_response.status_code, StatusCode::NOT_FOUND);
591 }
592
593 #[test]
594 fn test_error_response_task_not_cancelable() {
595 use serde_json::json;
596 let err = A2aError::TaskNotCancelable("Cannot cancel completed task".to_string());
597 let err_response = A2aErrorResponse::from_error(err, json!(1));
598 assert_eq!(err_response.status_code, StatusCode::UNPROCESSABLE_ENTITY);
599 }
600
601 #[test]
602 fn test_error_response_invalid_request() {
603 use serde_json::json;
604 let err_response = A2aErrorResponse::invalid_request("Invalid JSON", json!(1));
605 assert_eq!(err_response.status_code, StatusCode::BAD_REQUEST);
606 }
607
608 #[tokio::test]
609 async fn test_server_state_with_broadcast() {
610 let state = A2aServerState::vtcode_default("http://localhost:8080");
611
612 let mut rx = state.event_tx.subscribe();
614
615 let test_event = StreamingEvent::Message {
617 message: super::super::types::Message::agent_text("Test"),
618 context_id: Some("test".to_string()),
619 kind: "streaming-response".to_string(),
620 r#final: false,
621 };
622
623 let _ignored = state.event_tx.send(test_event.clone()).expect("send event");
624
625 let received = rx.recv().await.expect("receive event");
627 assert!(!received.is_final());
628 }
629
630 #[test]
631 fn test_server_state_debug_redacts_auth_token() {
632 let state = A2aServerState::new_with_auth_token(
633 TaskManager::new(),
634 AgentCard::vtcode_default("http://localhost:8080"),
635 "test-token",
636 )
637 .expect("valid auth token");
638
639 let debug = format!("{state:?}");
640 assert!(!debug.contains("test-token"));
641 assert!(debug.contains("REDACTED"));
642 }
643
644 #[test]
645 fn test_server_state_rejects_invalid_auth_token() {
646 let result = A2aServerState::new_with_auth_token(
647 TaskManager::new(),
648 AgentCard::vtcode_default("http://localhost:8080"),
649 " ",
650 );
651
652 assert!(result.is_err());
653 }
654
655 #[tokio::test]
656 async fn test_agent_card_is_public_and_advertises_authentication() {
657 use axum::{body::Body, http::Request};
658 use tower::ServiceExt;
659
660 let state = A2aServerState::new_with_auth_token(
661 TaskManager::new(),
662 AgentCard::vtcode_default("http://localhost:8080"),
663 "test-token",
664 )
665 .expect("valid auth token");
666 let app = create_router(state);
667
668 let response = app
669 .oneshot(
670 Request::builder()
671 .uri("/.well-known/agent-card.json")
672 .body(Body::empty())
673 .expect("build request"),
674 )
675 .await
676 .expect("receive response");
677
678 assert_eq!(response.status(), StatusCode::OK);
679 let body = axum::body::to_bytes(response.into_body(), usize::MAX).await.expect("read body");
680 let card: Value = serde_json::from_slice(&body).expect("parse card");
681 assert_eq!(card["securitySchemes"]["bearerAuth"]["scheme"], "bearer");
682 assert_eq!(card["security"][0]["bearerAuth"], json!([]));
683 }
684
685 #[tokio::test]
686 async fn test_rpc_requires_bearer_auth_before_json_parsing() {
687 use axum::{body::Body, http::Request};
688 use tower::ServiceExt;
689
690 let state = A2aServerState::new_with_auth_token(
691 TaskManager::new(),
692 AgentCard::vtcode_default("http://localhost:8080"),
693 "test-token",
694 )
695 .expect("valid auth token");
696 let token = state.auth_token().to_string();
697 let app = create_router(state);
698
699 let unauthorized = app
700 .clone()
701 .oneshot(
702 Request::builder()
703 .method("POST")
704 .uri("/a2a")
705 .header("content-type", "application/json")
706 .body(Body::from("not-json"))
707 .expect("build request"),
708 )
709 .await
710 .expect("receive response");
711 assert_eq!(unauthorized.status(), StatusCode::UNAUTHORIZED);
712 assert_eq!(unauthorized.headers()[WWW_AUTHENTICATE], "Bearer");
713
714 let wrong_token = app
715 .clone()
716 .oneshot(
717 Request::builder()
718 .method("POST")
719 .uri("/a2a")
720 .header("content-type", "application/json")
721 .header(AUTHORIZATION, "Bearer wrong-token")
722 .body(Body::from("{}"))
723 .expect("build request"),
724 )
725 .await
726 .expect("receive response");
727 assert_eq!(wrong_token.status(), StatusCode::UNAUTHORIZED);
728
729 let stream_unauthorized = app
730 .clone()
731 .oneshot(
732 Request::builder()
733 .method("POST")
734 .uri("/a2a/stream")
735 .header("content-type", "application/json")
736 .body(Body::from("not-json"))
737 .expect("build request"),
738 )
739 .await
740 .expect("receive response");
741 assert_eq!(stream_unauthorized.status(), StatusCode::UNAUTHORIZED);
742
743 let authenticated = app
744 .oneshot(
745 Request::builder()
746 .method("POST")
747 .uri("/a2a")
748 .header("content-type", "application/json")
749 .header(AUTHORIZATION, format!("Bearer {token}"))
750 .body(Body::from(r#"{"jsonrpc":"2.0","method":"unknown","id":1}"#))
751 .expect("build request"),
752 )
753 .await
754 .expect("receive response");
755 assert_eq!(authenticated.status(), StatusCode::NOT_FOUND);
756 }
757
758 #[tokio::test]
759 async fn test_duplicate_authorization_headers_are_rejected() {
760 use axum::{body::Body, http::Request};
761 use tower::ServiceExt;
762
763 let state = A2aServerState::new_with_auth_token(
764 TaskManager::new(),
765 AgentCard::vtcode_default("http://localhost:8080"),
766 "test-token",
767 )
768 .expect("valid auth token");
769 let app = create_router(state);
770 let mut request = Request::builder()
771 .method("POST")
772 .uri("/a2a")
773 .header("content-type", "application/json")
774 .body(Body::from("{}"))
775 .expect("build request");
776 let _ = request
777 .headers_mut()
778 .append(AUTHORIZATION, HeaderValue::from_static("Bearer test-token"));
779 let _ = request
780 .headers_mut()
781 .append(AUTHORIZATION, HeaderValue::from_static("Bearer test-token"));
782
783 let response = app.oneshot(request).await.expect("receive response");
784 assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
785 }
786
787 #[tokio::test]
788 async fn test_rpc_does_not_allow_cross_origin_cors_access() {
789 use axum::{body::Body, http::Request};
790 use tower::ServiceExt;
791
792 let state = A2aServerState::new_with_auth_token(
793 TaskManager::new(),
794 AgentCard::vtcode_default("http://localhost:8080"),
795 "test-token",
796 )
797 .expect("valid auth token");
798 let app = create_router(state);
799
800 let response = app
801 .oneshot(
802 Request::builder()
803 .method("POST")
804 .uri("/a2a")
805 .header("origin", "https://attacker.example")
806 .body(Body::from("{}"))
807 .expect("build request"),
808 )
809 .await
810 .expect("receive response");
811
812 assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
813 assert!(response.headers().get("access-control-allow-origin").is_none());
814 }
815
816 #[tokio::test]
817 async fn test_task_history_limits_are_applied_by_rpc_handlers() {
818 use axum::{body::Body, http::Request};
819 use tower::ServiceExt;
820
821 let task_manager = TaskManager::new();
822 let task = task_manager.create_task(None).await;
823 drop(
824 task_manager
825 .add_message(&task.id, crate::types::Message::user_text("private message"))
826 .await
827 .expect("add message"),
828 );
829 let state = A2aServerState::new_with_auth_token(
830 task_manager,
831 AgentCard::vtcode_default("http://localhost:8080"),
832 "test-token",
833 )
834 .expect("valid auth token");
835 let app = create_router(state);
836
837 let request_body = |params| {
838 serde_json::to_vec(&json!({
839 "jsonrpc": "2.0",
840 "method": METHOD_TASKS_GET,
841 "params": params,
842 "id": 1,
843 }))
844 .expect("serialize request")
845 };
846 let response = app
847 .clone()
848 .oneshot(
849 Request::builder()
850 .method("POST")
851 .uri("/a2a")
852 .header("content-type", "application/json")
853 .header(AUTHORIZATION, "Bearer test-token")
854 .body(Body::from(request_body(json!({"id": task.id}))))
855 .expect("build request"),
856 )
857 .await
858 .expect("receive response");
859 assert_eq!(response.status(), StatusCode::OK);
860 let body = axum::body::to_bytes(response.into_body(), usize::MAX).await.expect("read body");
861 let body: Value = serde_json::from_slice(&body).expect("parse response");
862 assert!(body["result"].get("history").is_none());
863
864 let response = app
865 .clone()
866 .oneshot(
867 Request::builder()
868 .method("POST")
869 .uri("/a2a")
870 .header("content-type", "application/json")
871 .header(AUTHORIZATION, "Bearer test-token")
872 .body(Body::from(request_body(json!({"id": task.id, "historyLength": 1}))))
873 .expect("build request"),
874 )
875 .await
876 .expect("receive response");
877 assert_eq!(response.status(), StatusCode::OK);
878 let body = axum::body::to_bytes(response.into_body(), usize::MAX).await.expect("read body");
879 let body: Value = serde_json::from_slice(&body).expect("parse response");
880 assert_eq!(body["result"]["history"].as_array().map(Vec::len), Some(1));
881 }
882
883 #[tokio::test]
884 async fn test_malformed_task_list_params_are_rejected() {
885 use axum::{body::Body, http::Request};
886 use tower::ServiceExt;
887
888 let state = A2aServerState::new_with_auth_token(
889 TaskManager::new(),
890 AgentCard::vtcode_default("http://localhost:8080"),
891 "test-token",
892 )
893 .expect("valid auth token");
894 let app = create_router(state);
895 let response = app
896 .oneshot(
897 Request::builder()
898 .method("POST")
899 .uri("/a2a")
900 .header("content-type", "application/json")
901 .header(AUTHORIZATION, "Bearer test-token")
902 .body(Body::from(r#"{"jsonrpc":"2.0","method":"tasks/list","params":[],"id":1}"#))
903 .expect("build request"),
904 )
905 .await
906 .expect("receive response");
907
908 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
909 }
910}