use crate::auth::RequestContext;
use crate::envelope::{codes, A2aHeaders, A2aMethod, JsonRpcError, JsonRpcResponse};
use crate::error::{A2aBuilderError, A2aError};
use crate::handler::A2aHandler;
use crate::server::{A2aDispatcher, TaskEvent, TaskEventStream};
use crate::task_store::{A2aTaskStore, DEFAULT_BUCKET};
use axum::{
body::Bytes,
extract::{DefaultBodyLimit, State},
http::{header, HeaderMap, StatusCode},
response::{sse::Event, IntoResponse, Response, Sse},
routing::post,
Json, Router,
};
use futures::StreamExt;
use klieo_auth_common::Authenticator;
use klieo_core::{KvStore, Pubsub};
use serde_json::Value;
use std::convert::Infallible;
use std::net::SocketAddr;
use std::sync::Arc;
use tokio_util::sync::CancellationToken;
use tracing::warn;
use tracing_opentelemetry::OpenTelemetrySpanExt as _;
const MAX_BODY_BYTES: usize = 1 << 20; const CANCEL_SUBJECT_PREFIX: &str = "klieo.a2a.cancel.";
pub struct A2aHttpServer {
pub(crate) dispatcher: Arc<A2aDispatcher>,
pub(crate) task_store: Arc<A2aTaskStore>,
pub(crate) parent_cancel: CancellationToken,
pub(crate) resume_buffer: Arc<dyn klieo_core::resume::ResumeBuffer>,
allow_public_bind: bool,
_kv_reaper: Option<klieo_core::KvReaperHandle>,
}
impl A2aHttpServer {
pub fn builder() -> A2aHttpServerBuilder {
A2aHttpServerBuilder::default()
}
pub fn new(
dispatcher: Arc<A2aDispatcher>,
task_store: Arc<A2aTaskStore>,
parent_cancel: CancellationToken,
) -> Self {
Self {
dispatcher,
task_store,
parent_cancel,
resume_buffer: Arc::new(klieo_core::resume::NoopResumeBuffer),
allow_public_bind: false,
_kv_reaper: None,
}
}
#[must_use]
pub fn with_resume_buffer(mut self, buffer: Arc<dyn klieo_core::resume::ResumeBuffer>) -> Self {
self.resume_buffer = buffer.clone();
self._kv_reaper = self.spawn_kv_reaper_if_configured(buffer);
self
}
fn spawn_kv_reaper_if_configured(
&self,
buffer: Arc<dyn klieo_core::resume::ResumeBuffer>,
) -> Option<klieo_core::KvReaperHandle> {
let interval = self.dispatcher.kv_reaper_interval()?;
let leader_registry = self.dispatcher.leader_registry()?;
let kv = leader_registry.kv().clone();
let mut buckets = vec![leader_registry.bucket().to_string()];
if let Some(ownership) = self.dispatcher.ownership_registry() {
buckets.push(ownership.bucket().to_string());
}
Some(klieo_core::spawn_kv_reaper(kv, buffer, buckets, interval))
}
pub fn allow_public_bind(mut self) -> Self {
self.allow_public_bind = true;
self
}
pub fn task_store(&self) -> &Arc<A2aTaskStore> {
&self.task_store
}
pub fn router(self: &Arc<Self>) -> Router {
Router::new()
.route("/a2a", post(post_a2a))
.layer(DefaultBodyLimit::max(MAX_BODY_BYTES))
.route(
"/.well-known/agent-card.json",
axum::routing::get(get_agent_card),
)
.with_state(self.clone())
}
pub async fn serve_http(self: Arc<Self>, addr: SocketAddr) -> Result<(), A2aError> {
if !addr.ip().is_loopback()
&& self.dispatcher.authenticator().allows_anonymous()
&& !self.allow_public_bind
{
tracing::error!(
target: "a2a",
%addr,
"refusing non-loopback bind with anonymous authenticator",
);
return Err(A2aError::Misconfigured(format!(
"refusing to bind non-loopback address {addr} with anonymous authenticator; \
call `allow_public_bind()` to override \
(only safe behind an auth-enforcing reverse proxy)",
)));
}
let cancel = self.parent_cancel.clone();
let listener = tokio::net::TcpListener::bind(addr)
.await
.map_err(|e| A2aError::Server(e.to_string()))?;
let router = self.router();
axum::serve(listener, router)
.with_graceful_shutdown(async move { cancel.cancelled().await })
.await
.map_err(|e| A2aError::Server(e.to_string()))?;
Ok(())
}
}
#[derive(Default)]
pub struct A2aHttpServerBuilder {
handler: Option<Arc<dyn A2aHandler>>,
authenticator: Option<Arc<dyn Authenticator>>,
kv: Option<Arc<dyn KvStore>>,
bucket: Option<String>,
cancel: Option<CancellationToken>,
pubsub: Option<Arc<dyn Pubsub>>,
}
impl A2aHttpServerBuilder {
pub fn handler(mut self, handler: Arc<dyn A2aHandler>) -> Self {
self.handler = Some(handler);
self
}
pub fn authenticator(mut self, authenticator: Arc<dyn Authenticator>) -> Self {
self.authenticator = Some(authenticator);
self
}
pub fn kv(mut self, kv: Arc<dyn KvStore>) -> Self {
self.kv = Some(kv);
self
}
pub fn bucket(mut self, bucket: String) -> Self {
self.bucket = Some(bucket);
self
}
pub fn cancel(mut self, cancel: CancellationToken) -> Self {
self.cancel = Some(cancel);
self
}
pub fn pubsub(mut self, pubsub: Arc<dyn Pubsub>) -> Self {
self.pubsub = Some(pubsub);
self
}
pub fn build(self) -> Result<A2aHttpServer, A2aBuilderError> {
let handler = self.handler.ok_or(A2aBuilderError::MissingHandler)?;
let authenticator = self
.authenticator
.ok_or(A2aBuilderError::MissingAuthenticator)?;
let kv = self.kv.ok_or(A2aBuilderError::MissingKv)?;
let dispatcher_builder = A2aDispatcher::builder()
.handler(handler)
.authenticator(authenticator);
let dispatcher_builder = match self.pubsub {
Some(pubsub) => dispatcher_builder.pubsub(pubsub),
None => dispatcher_builder.with_in_process_pubsub(),
};
let dispatcher = dispatcher_builder.build_arc()?;
let bucket = self.bucket.unwrap_or_else(|| DEFAULT_BUCKET.to_string());
let task_store =
Arc::new(A2aTaskStore::new(kv, bucket).with_event_sink(dispatcher.event_sink()));
let cancel = self.cancel.unwrap_or_default();
Ok(A2aHttpServer::new(dispatcher, task_store, cancel))
}
}
#[tracing::instrument(
skip_all,
fields(
rpc.system = "klieo-a2a",
rpc.method = tracing::field::Empty,
http.request.method = "POST",
),
)]
async fn post_a2a(
State(server): State<Arc<A2aHttpServer>>,
headers: HeaderMap,
body: Bytes,
) -> Response {
let parent_cx = klieo_core::extract_traceparent(&klieo_headers_from_axum(&headers));
tracing::Span::current().set_parent(parent_cx);
if server.parent_cancel.is_cancelled() {
return shutdown_response();
}
if !content_type_is_json(&headers) {
return StatusCode::UNSUPPORTED_MEDIA_TYPE.into_response();
}
if headers.get_all(header::AUTHORIZATION).iter().count() >= 2 {
warn!(target: "a2a", "rejected request with duplicate Authorization header");
return (
StatusCode::BAD_REQUEST,
Json(error_envelope(
serde_json::Value::Null,
codes::INVALID_REQUEST,
"duplicate Authorization header",
)),
)
.into_response();
}
let raw: Value = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(e) => {
warn!(error = %e, "rejected malformed A2A JSON-RPC body");
return (
StatusCode::BAD_REQUEST,
Json(error_envelope(
Value::Null,
codes::PARSE_ERROR,
"malformed JSON-RPC body",
)),
)
.into_response();
}
};
let method = raw.get("method").and_then(|m| m.as_str()).unwrap_or("");
let req_id = raw.get("id").cloned().unwrap_or(Value::Null);
tracing::Span::current().record("rpc.method", method);
if is_streaming_method(method) {
dispatch_streaming(&server, headers, body, req_id, method.to_owned()).await
} else {
dispatch_json(&server, headers, body).await
}
}
async fn get_agent_card(State(server): State<Arc<A2aHttpServer>>, headers: HeaderMap) -> Response {
let a2a_headers = axum_headers_to_a2a(&headers);
let ctx = RequestContext::new(a2a_headers, None);
match server.dispatcher.handler().get_agent_card(&ctx).await {
Ok(card) => (StatusCode::OK, Json(card)).into_response(),
Err(A2aError::MethodNotFound(_)) => StatusCode::NOT_FOUND.into_response(),
Err(err) => {
warn!(target: "a2a", error = %err, "agent-card discovery handler failed");
StatusCode::INTERNAL_SERVER_ERROR.into_response()
}
}
}
async fn dispatch_json(server: &Arc<A2aHttpServer>, headers: HeaderMap, body: Bytes) -> Response {
let a2a_headers = axum_headers_to_a2a(&headers);
let resp = server.dispatcher.handle_request(a2a_headers, &body).await;
(StatusCode::OK, Json(resp)).into_response()
}
async fn dispatch_streaming(
server: &Arc<A2aHttpServer>,
headers: HeaderMap,
body: Bytes,
req_id: Value,
method: String,
) -> Response {
let a2a_headers = axum_headers_to_a2a(&headers);
let last_event_id = last_event_id_from(&headers);
let request_cancel = server.parent_cancel.child_token();
let task_id = task_id_from_body(&method, &body);
match server
.dispatcher
.handle_streaming(
a2a_headers,
&body,
&server.task_store,
request_cancel.clone(),
last_event_id,
server.resume_buffer.clone(),
)
.await
{
Ok(stream) => build_sse_response(
stream,
req_id,
request_cancel,
server.dispatcher.clone(),
server.task_store.clone(),
server.resume_buffer.clone(),
task_id,
),
Err(A2aError::ResumeBufferExpired { since_id }) => (
StatusCode::OK,
Json(error_envelope(
req_id,
codes::RESUME_BUFFER_EXPIRED,
&format!("resume window expired (since_id={since_id})"),
)),
)
.into_response(),
Err(A2aError::Unauthorized(_)) => {
(
StatusCode::OK,
Json(error_envelope(
req_id,
codes::UNAUTHENTICATED,
"Authentication required",
)),
)
.into_response()
}
Err(err) => {
crate::server::log_internal_before_wire_seam(&err, &method);
if let A2aError::LeaderDied { stream_id } = &err {
return (
StatusCode::OK,
Json(leader_died_envelope(req_id, &err.to_string(), stream_id)),
)
.into_response();
}
let (code, msg) = match &err {
A2aError::InvalidParams(m) => (codes::INVALID_PARAMS, m.clone()),
A2aError::MethodNotFound(m) => {
(codes::METHOD_NOT_FOUND, format!("method not found: {m}"))
}
_ => (codes::SERVER_ERROR, "internal server error".into()),
};
(StatusCode::OK, Json(error_envelope(req_id, code, &msg))).into_response()
}
}
}
struct CancelOnDrop<S> {
inner: S,
_guard: tokio_util::sync::DropGuard,
pubsub: Arc<dyn klieo_core::Pubsub>,
cancel_subject: String,
permits: Arc<tokio::sync::Semaphore>,
}
impl<S: futures::Stream + Unpin> futures::Stream for CancelOnDrop<S> {
type Item = S::Item;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<S::Item>> {
std::pin::Pin::new(&mut self.inner).poll_next(cx)
}
}
impl<S> Drop for CancelOnDrop<S> {
fn drop(&mut self) {
let mut trace_headers = klieo_core::Headers::default();
klieo_core::inject_traceparent(&mut trace_headers, &opentelemetry::Context::current());
klieo_core::cancel::spawn_drop_publish(
self.pubsub.clone(),
std::mem::take(&mut self.cancel_subject),
"a2a.cancel",
Some(self.permits.clone()),
trace_headers,
);
}
}
fn task_id_from_body(method: &str, body: &Bytes) -> String {
if !matches!(A2aMethod::from_str(method), Ok(A2aMethod::SubscribeToTask)) {
return String::new();
}
let parsed: Result<Value, _> = serde_json::from_slice(body);
parsed
.ok()
.and_then(|v| v.get("params")?.get("id")?.as_str().map(str::to_owned))
.unwrap_or_default()
}
fn build_sse_response(
stream: TaskEventStream,
req_id: Value,
cancel: CancellationToken,
dispatcher: Arc<A2aDispatcher>,
task_store: Arc<A2aTaskStore>,
resume_buffer: Arc<dyn klieo_core::resume::ResumeBuffer>,
task_id: String,
) -> Response {
if !task_id.is_empty() {
dispatcher
.cancel_registry()
.register(task_id.clone(), cancel.clone());
}
let registry_handle = dispatcher.cancel_registry().clone();
let pubsub = dispatcher.pubsub().clone();
let permits = dispatcher.publish_permits().clone();
let cancel_subject = if task_id.is_empty() {
String::new()
} else {
format!("{CANCEL_SUBJECT_PREFIX}{task_id}")
};
let deregistered =
klieo_core::cancel::RegistryDeregisterOnDrop::new(stream, registry_handle, task_id);
let mapped = deregistered.map(move |mut event: TaskEvent| {
if event.event_id == 0 {
event.event_id = task_store.next_event_id(&event.task_id);
}
let id = event.event_id;
let final_event = event.final_event;
let task_id = event.task_id.clone();
let payload = bytes::Bytes::from(serde_json::to_vec(&event).unwrap_or_default());
let buffer = resume_buffer.clone();
tokio::spawn(async move {
if let Err(e) = buffer.record(&task_id, id, payload).await {
warn!(
target: "a2a.resume",
task_id = %task_id,
id,
error = %e,
"resume buffer record failed",
);
}
if final_event {
if let Err(e) = buffer.close(&task_id).await {
warn!(
target: "a2a.resume",
task_id = %task_id,
error = %e,
"resume buffer close failed",
);
}
}
});
task_event_to_sse_frame(event, &req_id, id)
});
let guarded = CancelOnDrop {
inner: mapped,
_guard: cancel.drop_guard(),
pubsub,
cancel_subject,
permits,
};
(StatusCode::OK, Sse::new(guarded)).into_response()
}
fn task_event_to_sse_frame(
event: TaskEvent,
req_id: &Value,
seq: u64,
) -> Result<Event, Infallible> {
let payload = serde_json::json!({
"jsonrpc": "2.0",
"id": req_id,
"result": {
"task_id": event.task_id,
"status": event.status,
"message": event.message,
"final": event.final_event,
},
});
let data = match serde_json::to_string(&payload) {
Ok(s) => s,
Err(e) => {
warn!(
target: "a2a",
task_id = %event.task_id,
error = %e,
"sse frame serialise failed; emitting comment",
);
return Ok(Event::default().comment("serialise-fail"));
}
};
Ok(Event::default()
.event("task-update")
.id(seq.to_string())
.data(data))
}
fn shutdown_response() -> Response {
(
StatusCode::SERVICE_UNAVAILABLE,
Json(JsonRpcResponse {
jsonrpc: "2.0".into(),
id: Value::Null,
result: None,
error: Some(JsonRpcError {
code: codes::SERVER_ERROR,
message: "server shutting down".into(),
data: None,
}),
}),
)
.into_response()
}
fn content_type_is_json(headers: &HeaderMap) -> bool {
headers
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(|s| {
s.split(';')
.next()
.unwrap_or("")
.trim()
.eq_ignore_ascii_case("application/json")
})
.unwrap_or(false)
}
fn is_streaming_method(method: &str) -> bool {
matches!(
A2aMethod::from_str(method),
Ok(A2aMethod::SendStreamingMessage) | Ok(A2aMethod::SubscribeToTask)
)
}
fn last_event_id_from(headers: &HeaderMap) -> Option<u64> {
headers
.get("last-event-id")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.trim().parse::<u64>().ok())
}
fn klieo_headers_from_axum(headers: &HeaderMap) -> klieo_core::Headers {
let mut out = klieo_core::Headers::default();
if let Some(value) = headers.get("traceparent").and_then(|v| v.to_str().ok()) {
out.insert("traceparent".into(), value.to_string());
}
if let Some(value) = headers.get("tracestate").and_then(|v| v.to_str().ok()) {
out.insert("tracestate".into(), value.to_string());
}
out
}
fn axum_headers_to_a2a(headers: &HeaderMap) -> A2aHeaders {
let mut kheaders = klieo_core::Headers::default();
for (name, value) in headers.iter() {
if let Ok(v) = value.to_str() {
kheaders.insert(name.as_str().to_string(), v.to_string());
}
}
A2aHeaders::decode_from(&kheaders)
}
fn error_envelope(id: Value, code: i32, message: &str) -> Value {
serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"error": { "code": code, "message": message },
})
}
fn leader_died_envelope(id: Value, message: &str, stream_id: &str) -> Value {
serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"error": {
"code": codes::LEADER_DIED,
"message": message,
"data": { "stream_id": stream_id },
},
})
}