Skip to main content

vtcode_a2a/
server.rs

1//! A2A HTTP Server using axum
2//!
3//! Provides HTTP endpoints for the A2A Protocol, enabling VT Code to operate as an A2A agent.
4//! The server exposes:
5//! - Agent discovery via `/.well-known/agent-card.json`
6//! - RPC endpoints at `/a2a` for message sending and task management
7//! - Streaming endpoint at `/a2a/stream` for real-time updates via Server-Sent Events
8
9use 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/// Bearer token used to protect the A2A RPC and streaming endpoints.
45#[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        // Branch-free OR-reduction over iterators (no per-byte bounds checks).
75        // LLVM can auto-vectorize this shape; the indexed `.get(i)` form cannot
76        // prove in-bounds and stays scalar. All bytes are still visited so the
77        // comparison remains constant-time with no early exit.
78        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// ============================================================================
94// Server State
95// ============================================================================
96
97/// A2A Server State containing shared resources
98#[derive(Debug, Clone)]
99pub struct A2aServerState {
100    /// Task manager for handling task lifecycle
101    task_manager: Arc<TaskManager>,
102    /// Agent card for discovery
103    agent_card: Arc<AgentCard>,
104    /// Broadcast channel for streaming events
105    event_tx: Arc<tokio::sync::broadcast::Sender<StreamingEvent>>,
106    /// Webhook notifier for push notifications
107    webhook_notifier: Arc<WebhookNotifier>,
108    /// Bearer token required by the RPC and streaming endpoints
109    auth_token: AuthToken,
110}
111
112impl A2aServerState {
113    /// Create a new server state
114    pub fn new(task_manager: TaskManager, agent_card: AgentCard) -> Self {
115        Self::from_auth_token(task_manager, agent_card, AuthToken::generated())
116    }
117
118    /// Create a new server state with a caller-supplied bearer token.
119    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    /// Return the bearer token clients must use for protected endpoints.
140    pub fn auth_token(&self) -> &str {
141        self.auth_token.as_str()
142    }
143
144    /// Create a server state with default settings for VT Code
145    fn vtcode_default(base_url: impl Into<String>) -> Self {
146        Self::new(TaskManager::new(), AgentCard::vtcode_default(base_url))
147    }
148}
149
150// ============================================================================
151// Router Creation
152// ============================================================================
153
154/// Create the A2A HTTP router
155pub 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        // Cross-origin access is disabled by default. A deployment that needs
166        // browser clients must add an explicit origin allowlist at its edge.
167        .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
211// ============================================================================
212// Handlers
213// ============================================================================
214
215/// Get agent card for discovery
216async fn get_agent_card(State(state): State<A2aServerState>) -> Json<AgentCard> {
217    Json(state.agent_card.as_ref().clone())
218}
219
220/// Handle JSON-RPC requests
221async fn handle_rpc(
222    State(state): State<A2aServerState>,
223    Json(request): Json<JsonRpcRequest>,
224) -> Result<Json<JsonRpcResponse>, A2aErrorResponse> {
225    // Validate request
226    if request.jsonrpc != JSONRPC_VERSION {
227        return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id));
228    }
229
230    // Dispatch to method handler
231    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
250/// Spawn a best-effort webhook delivery task.
251///
252/// Documented detached: the delivery is bounded (10s client timeout, max 3
253/// retries with backoff in `WebhookNotifier`), self-terminating, and its
254/// outcome is observable via `tracing::warn` on failure. Detachment is safe
255/// because webhook delivery must never block the SSE stream or agent turn;
256/// see "Task Extent, Error Propagation, and Cancel-Safety" in
257/// `docs/guides/async-architecture.md`.
258fn 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
274/// Handle Server-Sent Events streaming
275async 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    // Parse params
285    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    // Create or get task
289    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    // Add initial message
297    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    // Subscribe to broadcast channel
306    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    // Create stream from broadcast receiver using async_stream
313    let stream = async_stream::stream! {
314        while let Ok(event) = rx.recv().await {
315            // Filter events for this task/context
316            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    // Start background task to process and emit events.
341    // Documented detached: bounded demo pipeline (~600ms of sleeps plus
342    // channel sends), terminated by completion, with outcomes fanned out over
343    // the broadcast channel; webhook legs use `spawn_webhook_delivery`.
344    let state_clone = state.clone();
345    let task_id_clone = task_id.clone();
346    drop(tokio::spawn(async move {
347        // Simulate agent processing
348        tokio::time::sleep(Duration::from_millis(100)).await;
349
350        // Update task to working
351        drop(
352            state_clone
353                .task_manager
354                .update_status(&task_id_clone, TaskState::Working, None)
355                .await,
356        );
357
358        // Send status update event
359        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        // Simulate generating a response message
376        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        // Complete the task
394        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        // Send final status event
403        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
427// ============================================================================
428// RPC Method Handlers
429// ============================================================================
430
431/// Handle message/send RPC method
432async 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    // Create or get task
437    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    // Add message to history
445    drop(state.task_manager.add_message(&task_id, params.message).await?);
446
447    // Update status to working
448    let task = state.task_manager.update_status(&task_id, TaskState::Working, None).await?;
449
450    // Return task as response
451    Ok(serde_json::to_value(task)?)
452}
453
454/// Handle tasks/pushNotificationConfig/set RPC method
455async 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
464/// Handle tasks/pushNotificationConfig/get RPC method
465async 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(&params.id).await;
470
471    Ok(serde_json::to_value(config)?)
472}
473
474/// Handle message/stream RPC method
475fn handle_message_stream<'a>(
476    state: &'a A2aServerState,
477    params: Option<Value>,
478    id: Value,
479) -> impl Future<Output = A2aResult<Value>> + 'a {
480    // Same as message_send for now, but would support streaming
481    handle_message_send(state, params, id)
482}
483
484/// Handle tasks/get RPC method
485async 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(&params.id, params.history_length.unwrap_or_default() as usize)
492        .await?;
493
494    Ok(serde_json::to_value(task)?)
495}
496
497/// Handle tasks/list RPC method
498async 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
510/// Handle tasks/cancel RPC method
511async 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(&params.id).await?;
516
517    Ok(serde_json::to_value(task)?)
518}
519
520// ============================================================================
521// Error Response Handler
522// ============================================================================
523
524/// A2A error response for Axum
525pub struct A2aErrorResponse {
526    response: Box<JsonRpcResponse>,
527    status_code: StatusCode,
528}
529
530impl A2aErrorResponse {
531    /// Create a new error response
532    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    /// Create an invalid request error response
540    fn invalid_request(message: &str, id: Value) -> Self {
541        Self::new(JsonRpcError::invalid_request(message), id, StatusCode::BAD_REQUEST)
542    }
543
544    /// Create a method not found error response
545    fn method_not_found(method: &str, id: Value) -> Self {
546        Self::new(JsonRpcError::method_not_found(method), id, StatusCode::NOT_FOUND)
547    }
548
549    /// Create an error response from an A2aError
550    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
577// ============================================================================
578// Server Startup
579// ============================================================================
580
581/// Run the A2A server
582pub 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        // Verify broadcast channel works
628        let mut rx = state.event_tx.subscribe();
629
630        // Send a test event
631        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        // Receive the event
641        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}