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        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// ============================================================================
88// Server State
89// ============================================================================
90
91/// A2A Server State containing shared resources
92#[derive(Debug, Clone)]
93pub struct A2aServerState {
94    /// Task manager for handling task lifecycle
95    task_manager: Arc<TaskManager>,
96    /// Agent card for discovery
97    agent_card: Arc<AgentCard>,
98    /// Broadcast channel for streaming events
99    event_tx: Arc<tokio::sync::broadcast::Sender<StreamingEvent>>,
100    /// Webhook notifier for push notifications
101    webhook_notifier: Arc<WebhookNotifier>,
102    /// Bearer token required by the RPC and streaming endpoints
103    auth_token: AuthToken,
104}
105
106impl A2aServerState {
107    /// Create a new server state
108    pub fn new(task_manager: TaskManager, agent_card: AgentCard) -> Self {
109        Self::from_auth_token(task_manager, agent_card, AuthToken::generated())
110    }
111
112    /// Create a new server state with a caller-supplied bearer token.
113    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    /// Return the bearer token clients must use for protected endpoints.
134    pub fn auth_token(&self) -> &str {
135        self.auth_token.as_str()
136    }
137
138    /// Create a server state with default settings for VT Code
139    fn vtcode_default(base_url: impl Into<String>) -> Self {
140        Self::new(TaskManager::new(), AgentCard::vtcode_default(base_url))
141    }
142}
143
144// ============================================================================
145// Router Creation
146// ============================================================================
147
148/// Create the A2A HTTP router
149pub 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        // Cross-origin access is disabled by default. A deployment that needs
160        // browser clients must add an explicit origin allowlist at its edge.
161        .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
205// ============================================================================
206// Handlers
207// ============================================================================
208
209/// Get agent card for discovery
210async fn get_agent_card(State(state): State<A2aServerState>) -> Json<AgentCard> {
211    Json(state.agent_card.as_ref().clone())
212}
213
214/// Handle JSON-RPC requests
215async fn handle_rpc(
216    State(state): State<A2aServerState>,
217    Json(request): Json<JsonRpcRequest>,
218) -> Result<Json<JsonRpcResponse>, A2aErrorResponse> {
219    // Validate request
220    if request.jsonrpc != JSONRPC_VERSION {
221        return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id));
222    }
223
224    // Dispatch to method handler
225    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
244/// Spawn a best-effort webhook delivery task.
245///
246/// Documented detached: the delivery is bounded (10s client timeout, max 3
247/// retries with backoff in `WebhookNotifier`), self-terminating, and its
248/// outcome is observable via `tracing::warn` on failure. Detachment is safe
249/// because webhook delivery must never block the SSE stream or agent turn;
250/// see "Task Extent, Error Propagation, and Cancel-Safety" in
251/// `docs/guides/async-architecture.md`.
252fn 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
268/// Handle Server-Sent Events streaming
269async 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    // Parse params
279    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    // Create or get task
283    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    // Add initial message
291    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    // Subscribe to broadcast channel
300    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    // Create stream from broadcast receiver using async_stream
307    let stream = async_stream::stream! {
308        while let Ok(event) = rx.recv().await {
309            // Filter events for this task/context
310            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    // Start background task to process and emit events.
335    // Documented detached: bounded demo pipeline (~600ms of sleeps plus
336    // channel sends), terminated by completion, with outcomes fanned out over
337    // the broadcast channel; webhook legs use `spawn_webhook_delivery`.
338    let state_clone = state.clone();
339    let task_id_clone = task_id.clone();
340    drop(tokio::spawn(async move {
341        // Simulate agent processing
342        tokio::time::sleep(Duration::from_millis(100)).await;
343
344        // Update task to working
345        drop(
346            state_clone
347                .task_manager
348                .update_status(&task_id_clone, TaskState::Working, None)
349                .await,
350        );
351
352        // Send status update event
353        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        // Simulate generating a response message
370        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        // Complete the task
388        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        // Send final status event
397        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
421// ============================================================================
422// RPC Method Handlers
423// ============================================================================
424
425/// Handle message/send RPC method
426async 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    // Create or get task
431    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    // Add message to history
439    drop(state.task_manager.add_message(&task_id, params.message).await?);
440
441    // Update status to working
442    let task = state.task_manager.update_status(&task_id, TaskState::Working, None).await?;
443
444    // Return task as response
445    Ok(serde_json::to_value(task)?)
446}
447
448/// Handle tasks/pushNotificationConfig/set RPC method
449async 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
458/// Handle tasks/pushNotificationConfig/get RPC method
459async 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(&params.id).await;
464
465    Ok(serde_json::to_value(config)?)
466}
467
468/// Handle message/stream RPC method
469fn handle_message_stream<'a>(
470    state: &'a A2aServerState,
471    params: Option<Value>,
472    id: Value,
473) -> impl Future<Output = A2aResult<Value>> + 'a {
474    // Same as message_send for now, but would support streaming
475    handle_message_send(state, params, id)
476}
477
478/// Handle tasks/get RPC method
479async 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(&params.id, params.history_length.unwrap_or_default() as usize)
486        .await?;
487
488    Ok(serde_json::to_value(task)?)
489}
490
491/// Handle tasks/list RPC method
492async 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
504/// Handle tasks/cancel RPC method
505async 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(&params.id).await?;
510
511    Ok(serde_json::to_value(task)?)
512}
513
514// ============================================================================
515// Error Response Handler
516// ============================================================================
517
518/// A2A error response for Axum
519pub struct A2aErrorResponse {
520    response: JsonRpcResponse,
521    status_code: StatusCode,
522}
523
524impl A2aErrorResponse {
525    /// Create a new error response
526    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    /// Create an invalid request error response
534    fn invalid_request(message: &str, id: Value) -> Self {
535        Self::new(JsonRpcError::invalid_request(message), id, StatusCode::BAD_REQUEST)
536    }
537
538    /// Create a method not found error response
539    fn method_not_found(method: &str, id: Value) -> Self {
540        Self::new(JsonRpcError::method_not_found(method), id, StatusCode::NOT_FOUND)
541    }
542
543    /// Create an error response from an A2aError
544    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
571// ============================================================================
572// Server Startup
573// ============================================================================
574
575/// Run the A2A server
576pub 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        // Verify broadcast channel works
622        let mut rx = state.event_tx.subscribe();
623
624        // Send a test event
625        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        // Receive the event
635        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}