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