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