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::State,
23    http::StatusCode,
24    response::{
25        IntoResponse, Response,
26        sse::{Event, Sse},
27    },
28    routing::post,
29};
30use serde_json::{Value, json};
31use std::convert::Infallible;
32use std::future::Future;
33use std::net::SocketAddr;
34use std::sync::Arc;
35use std::time::Duration;
36use tower_http::cors::CorsLayer;
37
38// ============================================================================
39// Server State
40// ============================================================================
41
42/// A2A Server State containing shared resources
43#[derive(Debug, Clone)]
44pub struct A2aServerState {
45    /// Task manager for handling task lifecycle
46    task_manager: Arc<TaskManager>,
47    /// Agent card for discovery
48    agent_card: Arc<AgentCard>,
49    /// Broadcast channel for streaming events
50    event_tx: Arc<tokio::sync::broadcast::Sender<StreamingEvent>>,
51    /// Webhook notifier for push notifications
52    webhook_notifier: Arc<WebhookNotifier>,
53}
54
55impl A2aServerState {
56    /// Create a new server state
57    pub fn new(task_manager: TaskManager, agent_card: AgentCard) -> Self {
58        let (event_tx, _) = tokio::sync::broadcast::channel(100);
59        Self {
60            task_manager: Arc::new(task_manager),
61            agent_card: Arc::new(agent_card),
62            event_tx: Arc::new(event_tx),
63            webhook_notifier: Arc::new(WebhookNotifier::new()),
64        }
65    }
66
67    /// Create a server state with default settings for VT Code
68    fn vtcode_default(base_url: impl Into<String>) -> Self {
69        Self::new(TaskManager::new(), AgentCard::vtcode_default(base_url))
70    }
71}
72
73// ============================================================================
74// Router Creation
75// ============================================================================
76
77/// Create the A2A HTTP router
78pub fn create_router(state: A2aServerState) -> Router {
79    Router::new()
80        .route("/.well-known/agent-card.json", axum::routing::get(get_agent_card))
81        .route("/a2a", post(handle_rpc))
82        .route("/a2a/stream", post(handle_stream))
83        .with_state(state)
84        .layer(CorsLayer::permissive())
85}
86
87// ============================================================================
88// Handlers
89// ============================================================================
90
91/// Get agent card for discovery
92async fn get_agent_card(State(state): State<A2aServerState>) -> Json<AgentCard> {
93    Json(state.agent_card.as_ref().clone())
94}
95
96/// Handle JSON-RPC requests
97async fn handle_rpc(
98    State(state): State<A2aServerState>,
99    Json(request): Json<JsonRpcRequest>,
100) -> Result<Json<JsonRpcResponse>, A2aErrorResponse> {
101    // Validate request
102    if request.jsonrpc != JSONRPC_VERSION {
103        return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id));
104    }
105
106    // Dispatch to method handler
107    let result = match request.method.as_str() {
108        METHOD_MESSAGE_SEND => handle_message_send(&state, request.params, request.id.clone()).await,
109        METHOD_MESSAGE_STREAM => handle_message_stream(&state, request.params, request.id.clone()).await,
110        METHOD_TASKS_GET => handle_tasks_get(&state, request.params, request.id.clone()).await,
111        METHOD_TASKS_LIST => handle_tasks_list(&state, request.params, request.id.clone()).await,
112        METHOD_TASKS_CANCEL => handle_tasks_cancel(&state, request.params, request.id.clone()).await,
113        METHOD_TASKS_PUSH_CONFIG_SET => handle_push_config_set(&state, request.params, request.id.clone()).await,
114        METHOD_TASKS_PUSH_CONFIG_GET => handle_push_config_get(&state, request.params, request.id.clone()).await,
115        _ => {
116            return Err(A2aErrorResponse::method_not_found(&request.method, request.id));
117        }
118    };
119
120    match result {
121        Ok(result_value) => Ok(Json(JsonRpcResponse::success(result_value, request.id))),
122        Err(err) => Err(A2aErrorResponse::from_error(err, request.id)),
123    }
124}
125
126/// Handle Server-Sent Events streaming
127async fn handle_stream(State(state): State<A2aServerState>, Json(request): Json<JsonRpcRequest>) -> impl IntoResponse {
128    if request.jsonrpc != JSONRPC_VERSION {
129        return Err(A2aErrorResponse::invalid_request("Invalid JSON-RPC version", request.id.clone()));
130    }
131
132    if request.method != METHOD_MESSAGE_STREAM {
133        return Err(A2aErrorResponse::method_not_found(&request.method, request.id.clone()));
134    }
135
136    // Parse params
137    let params: MessageSendParams = serde_json::from_value(request.params.unwrap_or_default())
138        .map_err(|_e| A2aErrorResponse::invalid_request("Invalid message/stream params", request.id.clone()))?;
139
140    // Create or get task
141    let task_id = if let Some(task_id) = params.task_id.clone() {
142        task_id
143    } else {
144        let task = state.task_manager.create_task(params.context_id.clone()).await;
145        task.id.clone()
146    };
147
148    // Add initial message
149    drop(
150        state
151            .task_manager
152            .add_message(&task_id, params.message.clone())
153            .await
154            .map_err(|e| A2aErrorResponse::from_error(e, request.id.clone()))?,
155    );
156
157    // Subscribe to broadcast channel
158    let mut rx = state.event_tx.subscribe();
159    let task_id_clone = task_id.clone();
160    let context_id = params.context_id.clone();
161    let notifier = state.webhook_notifier.clone();
162    let task_manager = state.task_manager.clone();
163
164    // Create stream from broadcast receiver using async_stream
165    let stream = async_stream::stream! {
166        while let Ok(event) = rx.recv().await {
167            // Filter events for this task/context
168            let matches = match &event {
169                StreamingEvent::Message { context_id: ctx, .. } => {
170                    context_id.as_ref() == ctx.as_ref()
171                }
172                StreamingEvent::TaskStatus { task_id: tid, .. } => tid == &task_id_clone,
173                StreamingEvent::TaskArtifact { task_id: tid, .. } => tid == &task_id_clone,
174                _ => false,
175            };
176
177            if matches {
178                // Fire webhook asynchronously (best-effort)
179                let notifier = notifier.clone();
180                let task_manager = task_manager.clone();
181                let task_id_for_hook = task_id_clone.clone();
182                let event_for_hook = event.clone();
183                drop(tokio::spawn(async move {
184                    if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
185                        drop(notifier.send_event(&cfg, event_for_hook).await);
186                    }
187                }));
188
189                let is_final = event.is_final();
190                let json = serde_json::to_string(&SendStreamingMessageResponse { event })
191                    .unwrap_or_default();
192                yield Ok::<_, Infallible>(Event::default().data(json));
193
194                if is_final {
195                    break;
196                }
197            }
198        }
199    };
200
201    // Start background task to process and emit events
202    let state_clone = state.clone();
203    let task_id_clone = task_id.clone();
204    drop(tokio::spawn(async move {
205        // Simulate agent processing
206        tokio::time::sleep(Duration::from_millis(100)).await;
207
208        // Update task to working
209        drop(
210            state_clone
211                .task_manager
212                .update_status(&task_id_clone, TaskState::Working, None)
213                .await,
214        );
215
216        // Send status update event
217        let status_event = StreamingEvent::TaskStatus {
218            task_id: task_id_clone.clone(),
219            context_id: params.context_id.clone(),
220            status: crate::types::TaskStatus::new(TaskState::Working),
221            kind: "status-update".to_string(),
222            r#final: false,
223        };
224        drop(state_clone.event_tx.send(status_event.clone()));
225
226        // Fire webhook if configured
227        let notifier = state_clone.webhook_notifier.clone();
228        let task_manager = state_clone.task_manager.clone();
229        let task_id_for_hook = task_id_clone.clone();
230        drop(tokio::spawn(async move {
231            if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
232                drop(notifier.send_event(&cfg, status_event).await);
233            }
234        }));
235
236        // Simulate generating a response message
237        tokio::time::sleep(Duration::from_millis(200)).await;
238        let response_msg = crate::types::Message::agent_text("Processing your request...");
239        let message_event = StreamingEvent::Message {
240            message: response_msg,
241            context_id: params.context_id.clone(),
242            kind: "streaming-response".to_string(),
243            r#final: false,
244        };
245        drop(state_clone.event_tx.send(message_event.clone()));
246
247        // Fire webhook if configured
248        let notifier = state_clone.webhook_notifier.clone();
249        let task_manager = state_clone.task_manager.clone();
250        let task_id_for_hook = task_id_clone.clone();
251        drop(tokio::spawn(async move {
252            if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
253                drop(notifier.send_event(&cfg, message_event).await);
254            }
255        }));
256
257        // Complete the task
258        tokio::time::sleep(Duration::from_millis(300)).await;
259        drop(
260            state_clone
261                .task_manager
262                .update_status(&task_id_clone, TaskState::Completed, None)
263                .await,
264        );
265
266        // Send final status event
267        let final_status_event = StreamingEvent::TaskStatus {
268            task_id: task_id_clone,
269            context_id: params.context_id,
270            status: crate::types::TaskStatus::new(TaskState::Completed),
271            kind: "status-update".to_string(),
272            r#final: true,
273        };
274        drop(state_clone.event_tx.send(final_status_event.clone()));
275
276        // Fire webhook if configured
277        let notifier = state_clone.webhook_notifier.clone();
278        let task_manager = state_clone.task_manager.clone();
279        let task_id_for_hook = final_status_event.task_id().unwrap_or_default().to_string();
280        drop(tokio::spawn(async move {
281            if let Some(cfg) = task_manager.get_webhook_config(&task_id_for_hook).await {
282                drop(notifier.send_event(&cfg, final_status_event).await);
283            }
284        }));
285    }));
286
287    Ok(Sse::new(Box::pin(stream)).keep_alive(
288        axum::response::sse::KeepAlive::new()
289            .interval(Duration::from_secs(15))
290            .text("keep-alive"),
291    ))
292}
293
294// ============================================================================
295// RPC Method Handlers
296// ============================================================================
297
298/// Handle message/send RPC method
299async fn handle_message_send(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
300    let params: MessageSendParams = serde_json::from_value(params.unwrap_or_default())
301        .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid message/send params"))?;
302
303    // Create or get task
304    let task_id = if let Some(task_id) = params.task_id {
305        task_id
306    } else {
307        let task = state.task_manager.create_task(params.context_id).await;
308        task.id.clone()
309    };
310
311    // Add message to history
312    drop(state.task_manager.add_message(&task_id, params.message).await?);
313
314    // Update status to working
315    let task = state.task_manager.update_status(&task_id, TaskState::Working, None).await?;
316
317    // Return task as response
318    Ok(serde_json::to_value(task)?)
319}
320
321/// Handle tasks/pushNotificationConfig/set RPC method
322async fn handle_push_config_set(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
323    let config: crate::rpc::TaskPushNotificationConfig = serde_json::from_value(params.unwrap_or_default())
324        .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/set params"))?;
325
326    state.task_manager.set_webhook_config(config).await?;
327
328    Ok(json!({ "success": true }))
329}
330
331/// Handle tasks/pushNotificationConfig/get RPC method
332async fn handle_push_config_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
333    let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
334        .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid pushNotificationConfig/get params"))?;
335
336    let config = state.task_manager.get_webhook_config(&params.id).await;
337
338    Ok(serde_json::to_value(config)?)
339}
340
341/// Handle message/stream RPC method
342fn handle_message_stream<'a>(
343    state: &'a A2aServerState,
344    params: Option<Value>,
345    id: Value,
346) -> impl Future<Output = A2aResult<Value>> + 'a {
347    // Same as message_send for now, but would support streaming
348    handle_message_send(state, params, id)
349}
350
351/// Handle tasks/get RPC method
352async fn handle_tasks_get(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
353    let params: TaskQueryParams = serde_json::from_value(params.unwrap_or_default())
354        .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/get params"))?;
355
356    let task = state.task_manager.get_task_or_error(&params.id).await?;
357
358    Ok(serde_json::to_value(task)?)
359}
360
361/// Handle tasks/list RPC method
362async fn handle_tasks_list(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
363    let params: ListTasksParams = serde_json::from_value(params.unwrap_or_default()).unwrap_or_default();
364
365    let result = state.task_manager.list_tasks(params).await;
366
367    Ok(serde_json::to_value(result)?)
368}
369
370/// Handle tasks/cancel RPC method
371async fn handle_tasks_cancel(state: &A2aServerState, params: Option<Value>, _id: Value) -> A2aResult<Value> {
372    let params: TaskIdParams = serde_json::from_value(params.unwrap_or_default())
373        .map_err(|_e| A2aError::rpc(A2aErrorCode::InvalidParams, "Invalid tasks/cancel params"))?;
374
375    let task = state.task_manager.cancel_task(&params.id).await?;
376
377    Ok(serde_json::to_value(task)?)
378}
379
380// ============================================================================
381// Error Response Handler
382// ============================================================================
383
384/// A2A error response for Axum
385pub struct A2aErrorResponse {
386    response: JsonRpcResponse,
387    status_code: StatusCode,
388}
389
390impl A2aErrorResponse {
391    /// Create a new error response
392    fn new(error: JsonRpcError, id: Value, status_code: StatusCode) -> Self {
393        Self {
394            response: JsonRpcResponse::error(error, id),
395            status_code,
396        }
397    }
398
399    /// Create an invalid request error response
400    fn invalid_request(message: &str, id: Value) -> Self {
401        Self::new(JsonRpcError::invalid_request(message), id, StatusCode::BAD_REQUEST)
402    }
403
404    /// Create a method not found error response
405    fn method_not_found(method: &str, id: Value) -> Self {
406        Self::new(JsonRpcError::method_not_found(method), id, StatusCode::NOT_FOUND)
407    }
408
409    /// Create an error response from an A2aError
410    fn from_error(error: A2aError, id: Value) -> Self {
411        let code: i32 = error.code().into();
412        let message = error.to_string();
413        let status_code = match error {
414            A2aError::TaskNotFound(_) => StatusCode::NOT_FOUND,
415            A2aError::TaskNotCancelable(_) => StatusCode::UNPROCESSABLE_ENTITY,
416            A2aError::InvalidStateTransition { .. } => StatusCode::UNPROCESSABLE_ENTITY,
417            _ => StatusCode::INTERNAL_SERVER_ERROR,
418        };
419
420        Self::new(JsonRpcError::new(code, message), id, status_code)
421    }
422}
423
424impl IntoResponse for A2aErrorResponse {
425    fn into_response(self) -> Response {
426        (self.status_code, Json(self.response)).into_response()
427    }
428}
429
430// ============================================================================
431// Server Startup
432// ============================================================================
433
434/// Run the A2A server
435pub async fn run(state: A2aServerState, addr: SocketAddr) -> anyhow::Result<()> {
436    let listener = tokio::net::TcpListener::bind(addr).await?;
437    tracing::info!("A2A server listening on {}", addr);
438    axum::serve(listener, create_router(state))
439        .with_graceful_shutdown(crate::shutdown_signal_logged("A2A"))
440        .await?;
441    Ok(())
442}
443
444#[cfg(test)]
445mod tests {
446    use super::*;
447
448    #[test]
449    fn test_server_state_creation() {
450        let state = A2aServerState::vtcode_default("http://localhost:8080");
451        assert_eq!(state.agent_card.name, "vtcode-agent");
452    }
453
454    #[test]
455    fn test_error_response_task_not_found() {
456        use serde_json::json;
457        let err_response = A2aErrorResponse::from_error(A2aError::TaskNotFound("test-id".to_string()), json!(1));
458        assert_eq!(err_response.status_code, StatusCode::NOT_FOUND);
459    }
460
461    #[test]
462    fn test_error_response_task_not_cancelable() {
463        use serde_json::json;
464        let err = A2aError::TaskNotCancelable("Cannot cancel completed task".to_string());
465        let err_response = A2aErrorResponse::from_error(err, json!(1));
466        assert_eq!(err_response.status_code, StatusCode::UNPROCESSABLE_ENTITY);
467    }
468
469    #[test]
470    fn test_error_response_invalid_request() {
471        use serde_json::json;
472        let err_response = A2aErrorResponse::invalid_request("Invalid JSON", json!(1));
473        assert_eq!(err_response.status_code, StatusCode::BAD_REQUEST);
474    }
475
476    #[tokio::test]
477    async fn test_server_state_with_broadcast() {
478        let state = A2aServerState::vtcode_default("http://localhost:8080");
479
480        // Verify broadcast channel works
481        let mut rx = state.event_tx.subscribe();
482
483        // Send a test event
484        let test_event = StreamingEvent::Message {
485            message: super::super::types::Message::agent_text("Test"),
486            context_id: Some("test".to_string()),
487            kind: "streaming-response".to_string(),
488            r#final: false,
489        };
490
491        let _ignored = state.event_tx.send(test_event.clone()).expect("send event");
492
493        // Receive the event
494        let received = rx.recv().await.expect("receive event");
495        assert!(!received.is_final());
496    }
497}