use std::convert::Infallible;
use std::sync::Arc;
use axum::extract::State;
use axum::http::{header, HeaderMap, Method, StatusCode};
use axum::response::sse::{Event, Sse};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::{Json, Router};
use tokio::net::TcpListener;
use tokio_stream::wrappers::errors::BroadcastStreamRecvError;
use tokio_stream::wrappers::BroadcastStream;
use tokio_stream::StreamExt;
use tower_http::cors::{Any, CorsLayer};
use crate::protocol::{A2ARequest, A2AResponse, AgentCard};
use crate::server::A2AServer;
use crate::A2AError;
const CARD_PATH: &str = "/.well-known/agent-card.json";
const SSE_PATH: &str = "/events";
impl A2AServer {
pub async fn serve(self, port: u16) -> Result<(), A2AError> {
let addr = std::net::SocketAddr::from(([0, 0, 0, 0], port));
let listener = TcpListener::bind(addr)
.await
.map_err(|e| A2AError::Http(format!("Failed to bind {addr}: {e}")))?;
self.serve_on(listener).await
}
pub async fn serve_on(self, listener: TcpListener) -> Result<(), A2AError> {
let server = Arc::new(self);
axum::serve(listener, router(server))
.await
.map_err(|e| A2AError::Http(format!("Server error: {e}")))?;
Ok(())
}
}
fn router(server: Arc<A2AServer>) -> Router {
Router::new()
.route(CARD_PATH, get(get_agent_card))
.route("/", post(post_request))
.route(SSE_PATH, get(sse_stream))
.layer(
CorsLayer::new()
.allow_origin(Any)
.allow_methods([Method::GET, Method::POST])
.allow_headers(Any),
)
.with_state(server)
}
async fn get_agent_card(State(server): State<Arc<A2AServer>>) -> Json<AgentCard> {
Json(server.get_agent_card().clone())
}
async fn post_request(
State(server): State<Arc<A2AServer>>,
headers: HeaderMap,
body: String,
) -> Response {
let req: A2ARequest = match serde_json::from_str(&body) {
Ok(req) => req,
Err(e) => {
let resp = A2AResponse::error(0, -32700, format!("Invalid request: {e}"));
return (StatusCode::BAD_REQUEST, Json(resp)).into_response();
}
};
let bearer = headers
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "));
let resp = server.handle_a2a_request_authenticated(req, bearer).await;
(StatusCode::OK, Json(resp)).into_response()
}
async fn sse_stream(State(server): State<Arc<A2AServer>>) -> Response {
let Some(rx) = server.subscribe() else {
return (
StatusCode::NOT_FOUND,
"SSE not enabled (call A2AServer::with_streaming)",
)
.into_response();
};
let stream = BroadcastStream::new(rx).filter_map(|item| match item {
Ok(notification) => {
let data = serde_json::to_string(¬ification).ok()?;
Some(Ok::<_, Infallible>(
Event::default().event("task").data(data),
))
}
Err(BroadcastStreamRecvError::Lagged(_)) => {
Some(Ok(Event::default().event("reset").data("lagged")))
}
});
Sse::new(stream).into_response()
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
use std::time::Duration;
use lc_chains::base::{BaseChain, ChainError, ChainResult};
use serde_json::Value;
use crate::protocol::{A2AMessage, TaskStatus};
use crate::A2AClient;
struct EchoChain;
#[async_trait::async_trait]
impl BaseChain for EchoChain {
fn input_keys(&self) -> Vec<&str> {
vec!["input"]
}
fn output_keys(&self) -> Vec<&str> {
vec!["output"]
}
async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
let input = inputs
.get("input")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let mut out = HashMap::new();
out.insert("output".to_string(), Value::String(input));
Ok(out)
}
fn name(&self) -> &str {
"echo-chain"
}
}
async fn spawn(server: A2AServer) -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let handle = tokio::spawn(async move {
let _ = server.serve_on(listener).await;
});
(format!("http://127.0.0.1:{port}"), handle)
}
#[tokio::test]
async fn serve_exposes_agent_card() {
let server = A2AServer::new(Arc::new(EchoChain));
let (base, _handle) = spawn(server).await;
let client = A2AClient::new(base).unwrap();
let card = client.get_agent_card().await.unwrap();
assert_eq!(card.name, "echo-chain");
}
#[tokio::test]
async fn serve_dispatches_tasks_end_to_end() {
let server = A2AServer::new(Arc::new(EchoChain));
let (base, _handle) = spawn(server).await;
let client = A2AClient::new(base).unwrap();
let result = client
.send_task_and_wait(A2AMessage::user("hello"), Duration::from_secs(10))
.await
.unwrap();
assert_eq!(result.output, "hello");
}
#[tokio::test]
async fn serve_enforces_bearer_token() {
let server = A2AServer::new(Arc::new(EchoChain)).with_auth_token("secret-token");
let (base, _handle) = spawn(server).await;
let client = A2AClient::new(base.clone()).unwrap();
let err = client.send_task(A2AMessage::user("hi")).await.unwrap_err();
assert!(err.to_string().contains("Authentication required"));
let client = A2AClient::builder(base)
.bearer_token("secret-token")
.build()
.unwrap();
let result = client
.send_task_and_wait(A2AMessage::user("hi"), Duration::from_secs(10))
.await
.unwrap();
assert_eq!(result.output, "hi");
}
#[tokio::test]
async fn serve_streams_task_notifications() {
let server = A2AServer::new(Arc::new(EchoChain)).with_streaming(64);
let (base, _handle) = spawn(server).await;
let client = A2AClient::new(base.clone()).unwrap();
let mut stream = client
.send_task_streaming(&format!("{base}/events"), A2AMessage::user("hi"))
.await
.unwrap();
let mut saw_working = false;
let mut saw_completed = false;
while let Some(event) = stream.next().await {
let event = event.unwrap();
match event.status_value() {
Some(TaskStatus::Working) => saw_working = true,
Some(TaskStatus::Completed) => {
saw_completed = true;
break;
}
_ => {}
}
}
assert!(saw_working, "expected a working status-update");
assert!(saw_completed, "expected a completed status-update");
}
}