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
244fn spawn_webhook_delivery(
253 notifier: Arc<WebhookNotifier>,
254 task_manager: Arc<TaskManager>,
255 task_id: String,
256 event: StreamingEvent,
257) {
258 drop(tokio::spawn(async move {
259 let Some(cfg) = task_manager.get_webhook_config(&task_id).await else {
260 return;
261 };
262 if let Err(error) = notifier.send_event(&cfg, event).await {
263 tracing::warn!(task_id = %task_id, error = %error, "A2A webhook delivery failed");
264 }
265 }));
266}
267
268async fn handle_stream(State(state): State<A2aServerState>, Json(request): Json<JsonRpcRequest>) -> impl IntoResponse {
270 if request.jsonrpc != JSONRPC_VERSION {
271 return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id.clone()));
272 }
273
274 if request.method != METHOD_MESSAGE_STREAM {
275 return Err(A2aErrorResponse::method_not_found(&request.method, request.id.clone()));
276 }
277
278 let params: MessageSendParams = serde_json::from_value(request.params.unwrap_or_default())
280 .map_err(|_e| A2aErrorResponse::invalid_request("Invalid message/stream params", request.id.clone()))?;
281
282 let task_id = if let Some(task_id) = params.task_id.clone() {
284 task_id
285 } else {
286 let task = state.task_manager.create_task(params.context_id.clone()).await;
287 task.id.clone()
288 };
289
290 drop(
292 state
293 .task_manager
294 .add_message(&task_id, params.message.clone())
295 .await
296 .map_err(|e| A2aErrorResponse::from_error(e, request.id.clone()))?,
297 );
298
299 let mut rx = state.event_tx.subscribe();
301 let task_id_clone = task_id.clone();
302 let context_id = params.context_id.clone();
303 let notifier = state.webhook_notifier.clone();
304 let task_manager = state.task_manager.clone();
305
306 let stream = async_stream::stream! {
308 while let Ok(event) = rx.recv().await {
309 let matches = match &event {
311 StreamingEvent::Message { context_id: ctx, .. } => {
312 context_id.as_ref() == ctx.as_ref()
313 }
314 StreamingEvent::TaskStatus { task_id: tid, .. } => tid == &task_id_clone,
315 StreamingEvent::TaskArtifact { task_id: tid, .. } => tid == &task_id_clone,
316 _ => false,
317 };
318
319 if matches {
320 spawn_webhook_delivery(notifier.clone(), task_manager.clone(), task_id_clone.clone(), event.clone());
321
322 let is_final = event.is_final();
323 let json = serde_json::to_string(&SendStreamingMessageResponse { event })
324 .unwrap_or_default();
325 yield Ok::<_, Infallible>(Event::default().data(json));
326
327 if is_final {
328 break;
329 }
330 }
331 }
332 };
333
334 let state_clone = state.clone();
339 let task_id_clone = task_id.clone();
340 drop(tokio::spawn(async move {
341 tokio::time::sleep(Duration::from_millis(100)).await;
343
344 drop(
346 state_clone
347 .task_manager
348 .update_status(&task_id_clone, TaskState::Working, None)
349 .await,
350 );
351
352 let status_event = StreamingEvent::TaskStatus {
354 task_id: task_id_clone.clone(),
355 context_id: params.context_id.clone(),
356 status: crate::types::TaskStatus::new(TaskState::Working),
357 kind: "status-update".to_string(),
358 r#final: false,
359 };
360 drop(state_clone.event_tx.send(status_event.clone()));
361
362 spawn_webhook_delivery(
363 state_clone.webhook_notifier.clone(),
364 state_clone.task_manager.clone(),
365 task_id_clone.clone(),
366 status_event,
367 );
368
369 tokio::time::sleep(Duration::from_millis(200)).await;
371 let response_msg = crate::types::Message::agent_text("Processing your request...");
372 let message_event = StreamingEvent::Message {
373 message: response_msg,
374 context_id: params.context_id.clone(),
375 kind: "streaming-response".to_string(),
376 r#final: false,
377 };
378 drop(state_clone.event_tx.send(message_event.clone()));
379
380 spawn_webhook_delivery(
381 state_clone.webhook_notifier.clone(),
382 state_clone.task_manager.clone(),
383 task_id_clone.clone(),
384 message_event,
385 );
386
387 tokio::time::sleep(Duration::from_millis(300)).await;
389 drop(
390 state_clone
391 .task_manager
392 .update_status(&task_id_clone, TaskState::Completed, None)
393 .await,
394 );
395
396 let final_status_event = StreamingEvent::TaskStatus {
398 task_id: task_id_clone,
399 context_id: params.context_id,
400 status: crate::types::TaskStatus::new(TaskState::Completed),
401 kind: "status-update".to_string(),
402 r#final: true,
403 };
404 drop(state_clone.event_tx.send(final_status_event.clone()));
405
406 spawn_webhook_delivery(
407 state_clone.webhook_notifier.clone(),
408 state_clone.task_manager.clone(),
409 final_status_event.task_id().unwrap_or_default().to_string(),
410 final_status_event,
411 );
412 }));
413
414 Ok(Sse::new(Box::pin(stream)).keep_alive(
415 axum::response::sse::KeepAlive::new()
416 .interval(Duration::from_secs(15))
417 .text("keep-alive"),
418 ))
419}
420
421async fn handle_message_send(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
427 let params: MessageSendParams = serde_json::from_value(params.unwrap_or_default())
428 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid message/send params"))?;
429
430 let task_id = if let Some(task_id) = params.task_id {
432 task_id
433 } else {
434 let task = state.task_manager.create_task(params.context_id).await;
435 task.id.clone()
436 };
437
438 drop(state.task_manager.add_message(&task_id, params.message).await?);
440
441 let task = state.task_manager.update_status(&task_id, TaskState::Working, None).await?;
443
444 Ok(serde_json::to_value(task)?)
446}
447
448async fn handle_push_config_set(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
450 let config: crate::rpc::TaskPushNotificationConfig = serde_json::from_value(params.unwrap_or_default())
451 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/set params"))?;
452
453 state.task_manager.set_webhook_config(config).await?;
454
455 Ok(json!({ "success": true }))
456}
457
458async fn handle_push_config_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
460 let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
461 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/get params"))?;
462
463 let config = state.task_manager.get_webhook_config(¶ms.id).await;
464
465 Ok(serde_json::to_value(config)?)
466}
467
468fn handle_message_stream<'a>(
470 state: &'a A2aServerState,
471 params: Option<Value>,
472 id: Value,
473) -> impl Future<Output = A2aResult<Value>> + 'a {
474 handle_message_send(state, params, id)
476}
477
478async fn handle_tasks_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
480 let params: TaskQueryParams = serde_json::from_value(params.unwrap_or_default())
481 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/get params"))?;
482
483 let task = state
484 .task_manager
485 .get_task_or_error_with_history(¶ms.id, params.history_length.unwrap_or_default() as usize)
486 .await?;
487
488 Ok(serde_json::to_value(task)?)
489}
490
491async fn handle_tasks_list(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
493 let params = match params {
494 Some(params) => serde_json::from_value(params)
495 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/list params"))?,
496 None => ListTasksParams::default(),
497 };
498
499 let result = state.task_manager.list_tasks(params).await;
500
501 Ok(serde_json::to_value(result)?)
502}
503
504async fn handle_tasks_cancel(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
506 let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
507 .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/cancel params"))?;
508
509 let task = state.task_manager.cancel_task(¶ms.id).await?;
510
511 Ok(serde_json::to_value(task)?)
512}
513
514pub struct A2aErrorResponse {
520 response: JsonRpcResponse,
521 status_code: StatusCode,
522}
523
524impl A2aErrorResponse {
525 fn new(error: JsonRpcError, id: Value, status_code: StatusCode) -> Self {
527 Self {
528 response: JsonRpcResponse::error(error, id),
529 status_code,
530 }
531 }
532
533 fn invalid_request(message: &str, id: Value) -> Self {
535 Self::new(JsonRpcError::invalid_request(message), id, StatusCode::BAD_REQUEST)
536 }
537
538 fn method_not_found(method: &str, id: Value) -> Self {
540 Self::new(JsonRpcError::method_not_found(method), id, StatusCode::NOT_FOUND)
541 }
542
543 fn from_error(error: A2aError, id: Value) -> Self {
545 let code: i32 = error.code().into();
546 let message = error.to_string();
547 let status_code = match &error {
548 A2aError::RpcError { code, .. } => match code {
549 A2aErrorCode::InvalidRequest | A2aErrorCode::InvalidParams | A2aErrorCode::JsonParseError => {
550 StatusCode::BAD_REQUEST
551 }
552 A2aErrorCode::MethodNotFound => StatusCode::NOT_FOUND,
553 _ => StatusCode::INTERNAL_SERVER_ERROR,
554 },
555 A2aError::TaskNotFound(_) => StatusCode::NOT_FOUND,
556 A2aError::TaskNotCancelable(_) => StatusCode::UNPROCESSABLE_ENTITY,
557 A2aError::InvalidStateTransition { .. } => StatusCode::UNPROCESSABLE_ENTITY,
558 _ => StatusCode::INTERNAL_SERVER_ERROR,
559 };
560
561 Self::new(JsonRpcError::new(code, message), id, status_code)
562 }
563}
564
565impl IntoResponse for A2aErrorResponse {
566 fn into_response(self) -> Response {
567 (self.status_code, Json(self.response)).into_response()
568 }
569}
570
571pub async fn run(state: A2aServerState, addr: SocketAddr) -> anyhow::Result<()> {
577 let listener = tokio::net::TcpListener::bind(addr).await?;
578 tracing::info!("A2A server listening on {}", addr);
579 axum::serve(listener, create_router(state))
580 .with_graceful_shutdown(crate::shutdown_signal_logged("A2A"))
581 .await?;
582 Ok(())
583}
584
585#[cfg(test)]
586mod tests {
587 use super::*;
588
589 #[test]
590 fn test_server_state_creation() {
591 let state = A2aServerState::vtcode_default("http://localhost:8080");
592 assert_eq!(state.agent_card.name, "vtcode-agent");
593 }
594
595 #[test]
596 fn test_error_response_task_not_found() {
597 use serde_json::json;
598 let err_response = A2aErrorResponse::from_error(A2aError::TaskNotFound("test-id".to_string()), json!(1));
599 assert_eq!(err_response.status_code, StatusCode::NOT_FOUND);
600 }
601
602 #[test]
603 fn test_error_response_task_not_cancelable() {
604 use serde_json::json;
605 let err = A2aError::TaskNotCancelable("Cannot cancel completed task".to_string());
606 let err_response = A2aErrorResponse::from_error(err, json!(1));
607 assert_eq!(err_response.status_code, StatusCode::UNPROCESSABLE_ENTITY);
608 }
609
610 #[test]
611 fn test_error_response_invalid_request() {
612 use serde_json::json;
613 let err_response = A2aErrorResponse::invalid_request("Invalid JSON", json!(1));
614 assert_eq!(err_response.status_code, StatusCode::BAD_REQUEST);
615 }
616
617 #[tokio::test]
618 async fn test_server_state_with_broadcast() {
619 let state = A2aServerState::vtcode_default("http://localhost:8080");
620
621 let mut rx = state.event_tx.subscribe();
623
624 let test_event = StreamingEvent::Message {
626 message: super::super::types::Message::agent_text("Test"),
627 context_id: Some("test".to_string()),
628 kind: "streaming-response".to_string(),
629 r#final: false,
630 };
631
632 let _ignored = state.event_tx.send(test_event.clone()).expect("send event");
633
634 let received = rx.recv().await.expect("receive event");
636 assert!(!received.is_final());
637 }
638
639 #[test]
640 fn test_server_state_debug_redacts_auth_token() {
641 let state = A2aServerState::new_with_auth_token(
642 TaskManager::new(),
643 AgentCard::vtcode_default("http://localhost:8080"),
644 "test-token",
645 )
646 .expect("valid auth token");
647
648 let debug = format!("{state:?}");
649 assert!(!debug.contains("test-token"));
650 assert!(debug.contains("REDACTED"));
651 }
652
653 #[test]
654 fn test_server_state_rejects_invalid_auth_token() {
655 let result = A2aServerState::new_with_auth_token(
656 TaskManager::new(),
657 AgentCard::vtcode_default("http://localhost:8080"),
658 " ",
659 );
660
661 assert!(result.is_err());
662 }
663
664 #[tokio::test]
665 async fn test_agent_card_is_public_and_advertises_authentication() {
666 use axum::{body::Body, http::Request};
667 use tower::ServiceExt;
668
669 let state = A2aServerState::new_with_auth_token(
670 TaskManager::new(),
671 AgentCard::vtcode_default("http://localhost:8080"),
672 "test-token",
673 )
674 .expect("valid auth token");
675 let app = create_router(state);
676
677 let response = app
678 .oneshot(
679 Request::builder()
680 .uri("/.well-known/agent-card.json")
681 .body(Body::empty())
682 .expect("build request"),
683 )
684 .await
685 .expect("receive response");
686
687 assert_eq!(response.status(), StatusCode::OK);
688 let body = axum::body::to_bytes(response.into_body(), usize::MAX).await.expect("read body");
689 let card: Value = serde_json::from_slice(&body).expect("parse card");
690 assert_eq!(card["securitySchemes"]["bearerAuth"]["scheme"], "bearer");
691 assert_eq!(card["security"][0]["bearerAuth"], json!([]));
692 }
693
694 #[tokio::test]
695 async fn test_rpc_requires_bearer_auth_before_json_parsing() {
696 use axum::{body::Body, http::Request};
697 use tower::ServiceExt;
698
699 let state = A2aServerState::new_with_auth_token(
700 TaskManager::new(),
701 AgentCard::vtcode_default("http://localhost:8080"),
702 "test-token",
703 )
704 .expect("valid auth token");
705 let token = state.auth_token().to_string();
706 let app = create_router(state);
707
708 let unauthorized = app
709 .clone()
710 .oneshot(
711 Request::builder()
712 .method("POST")
713 .uri("/a2a")
714 .header("content-type", "application/json")
715 .body(Body::from("not-json"))
716 .expect("build request"),
717 )
718 .await
719 .expect("receive response");
720 assert_eq!(unauthorized.status(), StatusCode::UNAUTHORIZED);
721 assert_eq!(unauthorized.headers()[WWW_AUTHENTICATE], "Bearer");
722
723 let wrong_token = app
724 .clone()
725 .oneshot(
726 Request::builder()
727 .method("POST")
728 .uri("/a2a")
729 .header("content-type", "application/json")
730 .header(AUTHORIZATION, "Bearer wrong-token")
731 .body(Body::from("{}"))
732 .expect("build request"),
733 )
734 .await
735 .expect("receive response");
736 assert_eq!(wrong_token.status(), StatusCode::UNAUTHORIZED);
737
738 let stream_unauthorized = app
739 .clone()
740 .oneshot(
741 Request::builder()
742 .method("POST")
743 .uri("/a2a/stream")
744 .header("content-type", "application/json")
745 .body(Body::from("not-json"))
746 .expect("build request"),
747 )
748 .await
749 .expect("receive response");
750 assert_eq!(stream_unauthorized.status(), StatusCode::UNAUTHORIZED);
751
752 let authenticated = app
753 .oneshot(
754 Request::builder()
755 .method("POST")
756 .uri("/a2a")
757 .header("content-type", "application/json")
758 .header(AUTHORIZATION, format!("Bearer {token}"))
759 .body(Body::from(r#"{"jsonrpc":"2.0","method":"unknown","id":1}"#))
760 .expect("build request"),
761 )
762 .await
763 .expect("receive response");
764 assert_eq!(authenticated.status(), StatusCode::NOT_FOUND);
765 }
766
767 #[tokio::test]
768 async fn test_duplicate_authorization_headers_are_rejected() {
769 use axum::{body::Body, http::Request};
770 use tower::ServiceExt;
771
772 let state = A2aServerState::new_with_auth_token(
773 TaskManager::new(),
774 AgentCard::vtcode_default("http://localhost:8080"),
775 "test-token",
776 )
777 .expect("valid auth token");
778 let app = create_router(state);
779 let mut request = Request::builder()
780 .method("POST")
781 .uri("/a2a")
782 .header("content-type", "application/json")
783 .body(Body::from("{}"))
784 .expect("build request");
785 let _ = request
786 .headers_mut()
787 .append(AUTHORIZATION, HeaderValue::from_static("Bearer test-token"));
788 let _ = request
789 .headers_mut()
790 .append(AUTHORIZATION, HeaderValue::from_static("Bearer test-token"));
791
792 let response = app.oneshot(request).await.expect("receive response");
793 assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
794 }
795
796 #[tokio::test]
797 async fn test_rpc_does_not_allow_cross_origin_cors_access() {
798 use axum::{body::Body, http::Request};
799 use tower::ServiceExt;
800
801 let state = A2aServerState::new_with_auth_token(
802 TaskManager::new(),
803 AgentCard::vtcode_default("http://localhost:8080"),
804 "test-token",
805 )
806 .expect("valid auth token");
807 let app = create_router(state);
808
809 let response = app
810 .oneshot(
811 Request::builder()
812 .method("POST")
813 .uri("/a2a")
814 .header("origin", "https://attacker.example")
815 .body(Body::from("{}"))
816 .expect("build request"),
817 )
818 .await
819 .expect("receive response");
820
821 assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
822 assert!(response.headers().get("access-control-allow-origin").is_none());
823 }
824
825 #[tokio::test]
826 async fn test_task_history_limits_are_applied_by_rpc_handlers() {
827 use axum::{body::Body, http::Request};
828 use tower::ServiceExt;
829
830 let task_manager = TaskManager::new();
831 let task = task_manager.create_task(None).await;
832 drop(
833 task_manager
834 .add_message(&task.id, crate::types::Message::user_text("private message"))
835 .await
836 .expect("add message"),
837 );
838 let state = A2aServerState::new_with_auth_token(
839 task_manager,
840 AgentCard::vtcode_default("http://localhost:8080"),
841 "test-token",
842 )
843 .expect("valid auth token");
844 let app = create_router(state);
845
846 let request_body = |params| {
847 serde_json::to_vec(&json!({
848 "jsonrpc": "2.0",
849 "method": METHOD_TASKS_GET,
850 "params": params,
851 "id": 1,
852 }))
853 .expect("serialize request")
854 };
855 let response = app
856 .clone()
857 .oneshot(
858 Request::builder()
859 .method("POST")
860 .uri("/a2a")
861 .header("content-type", "application/json")
862 .header(AUTHORIZATION, "Bearer test-token")
863 .body(Body::from(request_body(json!({"id": task.id}))))
864 .expect("build request"),
865 )
866 .await
867 .expect("receive response");
868 assert_eq!(response.status(), StatusCode::OK);
869 let body = axum::body::to_bytes(response.into_body(), usize::MAX).await.expect("read body");
870 let body: Value = serde_json::from_slice(&body).expect("parse response");
871 assert!(body["result"].get("history").is_none());
872
873 let response = app
874 .clone()
875 .oneshot(
876 Request::builder()
877 .method("POST")
878 .uri("/a2a")
879 .header("content-type", "application/json")
880 .header(AUTHORIZATION, "Bearer test-token")
881 .body(Body::from(request_body(json!({"id": task.id, "historyLength": 1}))))
882 .expect("build request"),
883 )
884 .await
885 .expect("receive response");
886 assert_eq!(response.status(), StatusCode::OK);
887 let body = axum::body::to_bytes(response.into_body(), usize::MAX).await.expect("read body");
888 let body: Value = serde_json::from_slice(&body).expect("parse response");
889 assert_eq!(body["result"]["history"].as_array().map(Vec::len), Some(1));
890 }
891
892 #[tokio::test]
893 async fn test_malformed_task_list_params_are_rejected() {
894 use axum::{body::Body, http::Request};
895 use tower::ServiceExt;
896
897 let state = A2aServerState::new_with_auth_token(
898 TaskManager::new(),
899 AgentCard::vtcode_default("http://localhost:8080"),
900 "test-token",
901 )
902 .expect("valid auth token");
903 let app = create_router(state);
904 let response = app
905 .oneshot(
906 Request::builder()
907 .method("POST")
908 .uri("/a2a")
909 .header("content-type", "application/json")
910 .header(AUTHORIZATION, "Bearer test-token")
911 .body(Body::from(r#"{"jsonrpc":"2.0","method":"tasks/list","params":[],"id":1}"#))
912 .expect("build request"),
913 )
914 .await
915 .expect("receive response");
916
917 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
918 }
919}