use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use axum::Router;
use axum::body::{Body, Bytes};
use axum::extract::State;
use axum::http::{HeaderMap, StatusCode, header};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use futures_util::StreamExt;
use serde_json::{Value, json};
use tower::ServiceExt;
use crate::a2a::Principal;
use crate::a2a::ports::{self, RuntimePorts};
use crate::obs::log::Logger;
use crate::runtime::a2a_server::{A2aBridge, PairingState, SharedFeed};
pub struct Auth {
pub require_auth: bool,
pub server_bearer: Option<String>,
pub pairing: Option<Arc<PairingState>>,
}
pub struct Opts {
pub auth: Auth,
pub extra_origins: Vec<String>,
pub tls: Option<TlsConfigProvider>,
pub request_timeout: Duration,
pub stream_deadline: Duration,
}
pub struct Listener {
pub bound: String,
pub sink: Arc<ports::StreamSink>,
_runtime: tokio::runtime::Runtime,
}
struct App {
protocol: Router,
bridge: Arc<A2aBridge>,
auth: Auth,
extra_origins: Vec<String>,
stream_deadline: Duration,
log: Logger,
}
pub enum Bind {
Tcp(String),
Unix(String),
}
pub fn spawn(
bind: Bind,
opts: Opts,
bridge: Arc<A2aBridge>,
feed: Option<Arc<SharedFeed>>,
log: Logger,
) -> Result<Listener, String> {
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_all()
.thread_name("agentd-a2a")
.build()
.map_err(|e| format!("a2a runtime: {e}"))?;
let updates = Arc::new(a2a_rs::adapter::InMemoryStreamingHandler::new());
let sink = Arc::new(ports::StreamSink::new(
Arc::clone(&updates),
runtime.handle().clone(),
log.clone(),
));
let ports = RuntimePorts::new(Arc::clone(&bridge), Arc::clone(&updates));
let adapter = Arc::new(
a2a_rs::adapter::JsonRpcAdapter::with_handler(ports, CardFromRuntime(Arc::clone(&bridge)))
.with_streaming_handler(ports::SharedStreaming(updates)),
);
let app = Arc::new(App {
protocol: a2a_rs::adapter::jsonrpc_router(adapter),
bridge: Arc::clone(&bridge),
auth: opts.auth,
extra_origins: opts.extra_origins,
stream_deadline: opts.stream_deadline,
log: log.clone(),
});
let _ = feed;
let router = Router::new()
.route("/", post(rpc).options(preflight))
.route("/.well-known/agent-card.json", get(card))
.route("/.well-known/agent.json", get(card))
.with_state(Arc::clone(&app));
let bound;
match &bind {
Bind::Tcp(authority) => {
let listener = runtime
.block_on(tokio::net::TcpListener::bind(authority))
.map_err(|e| format!("a2a bind {authority}: {e}"))?;
bound = listener
.local_addr()
.map(|a| a.to_string())
.unwrap_or_else(|_| authority.clone());
let tls = opts.tls;
runtime.spawn(accept_loop(listener, router, tls, log));
}
Bind::Unix(path) => {
let _ = std::fs::remove_file(path);
let sock = std::path::Path::new(path);
let stage = sock
.parent()
.unwrap_or(std::path::Path::new("."))
.join(format!(".agentd-sock-{}", std::process::id()));
let _ = std::fs::remove_dir_all(&stage);
std::fs::create_dir(&stage)
.map_err(|e| format!("a2a stage {}: {e}", stage.display()))?;
std::fs::set_permissions(&stage, std::os::unix::fs::PermissionsExt::from_mode(0o700))
.map_err(|e| format!("a2a stage perms: {e}"))?;
let staged = stage.join("s");
let bound_listener = {
let _guard = runtime.enter();
tokio::net::UnixListener::bind(&staged)
};
let listener = match bound_listener {
Ok(l) => l,
Err(e) => {
let _ = std::fs::remove_dir_all(&stage);
return Err(format!("a2a bind {path}: {e}"));
}
};
let perms = std::fs::set_permissions(
&staged,
std::os::unix::fs::PermissionsExt::from_mode(0o600),
);
let published = perms.and_then(|()| std::fs::rename(&staged, sock));
let _ = std::fs::remove_dir_all(&stage);
published.map_err(|e| format!("a2a publish {path}: {e}"))?;
bound = format!("unix:{path}");
runtime.spawn(accept_loop_unix(listener, router, log));
}
}
Ok(Listener {
bound,
sink,
_runtime: runtime,
})
}
pub type TlsConfigProvider = Arc<dyn Fn() -> Arc<tokio_rustls::rustls::ServerConfig> + Send + Sync>;
async fn accept_loop(
listener: tokio::net::TcpListener,
router: Router,
tls: Option<TlsConfigProvider>,
log: Logger,
) {
loop {
let Ok((sock, peer)) = listener.accept().await else {
continue;
};
let router = router.clone();
let tls = tls.clone();
let log = log.clone();
tokio::spawn(async move {
match tls {
Some(provider) => {
let acceptor = tokio_rustls::TlsAcceptor::from(provider());
match acceptor.accept(sock).await {
Ok(stream) => {
let peer_id = peer_identity(stream.get_ref().1);
serve_conn(stream, router, peer_id, peer, log).await;
}
Err(e) => log.debug("a2a.tls", json!({"err": e.to_string()})),
}
}
None => serve_conn(sock, router, PeerId::default(), peer, log).await,
}
});
}
}
async fn accept_loop_unix(listener: tokio::net::UnixListener, router: Router, log: Logger) {
let me = unsafe { libc::geteuid() };
loop {
let Ok((sock, _)) = listener.accept().await else {
continue;
};
let uid = match sock.peer_cred() {
Ok(cred) => cred.uid(),
Err(e) => {
log.warn("a2a.unix.denied", json!({"err": e.to_string()}));
continue;
}
};
if uid != me && uid != 0 {
log.warn("a2a.unix.denied", json!({"uid": uid, "reason": "peer uid"}));
continue;
}
let router = router.clone();
let log = log.clone();
tokio::spawn(async move {
let peer: SocketAddr = "127.0.0.1:0".parse().expect("static addr");
serve_conn(sock, router, PeerId::default(), peer, log).await;
});
}
}
#[derive(Clone, Default, Debug)]
pub struct PeerId {
pub presented: bool,
pub subject: Option<String>,
pub sans: Vec<String>,
}
fn peer_identity(conn: &tokio_rustls::rustls::ServerConnection) -> PeerId {
let Some(chain) = conn.peer_certificates() else {
return PeerId::default();
};
let Some(leaf) = chain.first() else {
return PeerId {
presented: true,
..Default::default()
};
};
let id = crate::net::x509::parse(leaf.as_ref());
PeerId {
presented: true,
subject: id.subject_cn,
sans: id.sans,
}
}
async fn serve_conn<S>(stream: S, router: Router, peer_id: PeerId, peer: SocketAddr, log: Logger)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let router = router
.layer(axum::Extension(peer_id))
.layer(axum::Extension(Peer(peer)));
let svc = hyper_util::service::TowerToHyperService::new(
router.into_service::<hyper::body::Incoming>(),
);
let io = hyper_util::rt::TokioIo::new(stream);
if let Err(e) = hyper::server::conn::http1::Builder::new()
.serve_connection(io, svc)
.with_upgrades()
.await
{
log.debug("a2a.conn", json!({"err": e.to_string()}));
}
}
#[derive(Clone, Copy)]
struct Peer(SocketAddr);
struct CardFromRuntime(Arc<A2aBridge>);
#[async_trait::async_trait]
impl a2a_rs::services::AgentInfoProvider for CardFromRuntime {
async fn get_agent_card(&self) -> Result<a2a_rs::domain::AgentCard, a2a_rs::domain::A2AError> {
self.card("GetAgentCard", Principal::anonymous()).await
}
async fn get_authenticated_extended_card(
&self,
) -> Result<a2a_rs::domain::AgentCard, a2a_rs::domain::A2AError> {
self.card("GetExtendedAgentCard", ports::caller()).await
}
}
impl CardFromRuntime {
async fn card(
&self,
method: &'static str,
who: Principal,
) -> Result<a2a_rs::domain::AgentCard, a2a_rs::domain::A2AError> {
let bridge = Arc::clone(&self.0);
let v = tokio::task::spawn_blocking(move || bridge.call(method, json!({}), who))
.await
.map_err(|e| a2a_rs::domain::A2AError::Internal(e.to_string()))?;
if let Some(e) = v.get("_error") {
return Err(a2a_rs::domain::A2AError::UnsupportedOperation(
e.get("message")
.and_then(Value::as_str)
.unwrap_or("no extended card")
.to_string(),
));
}
serde_json::from_value(v).map_err(a2a_rs::domain::A2AError::JsonParse)
}
}
async fn preflight(State(app): State<Arc<App>>, headers: HeaderMap) -> Response {
let origin = headers
.get(header::ORIGIN)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if !origin_allowed(origin, &app.extra_origins) {
return (StatusCode::FORBIDDEN, "").into_response();
}
let wants_private_network = headers
.get("access-control-request-private-network")
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.eq_ignore_ascii_case("true"));
let mut resp = (
StatusCode::NO_CONTENT,
[
(header::ACCESS_CONTROL_ALLOW_ORIGIN, origin.to_string()),
(
header::ACCESS_CONTROL_ALLOW_METHODS,
"POST, GET, OPTIONS".to_string(),
),
(
header::ACCESS_CONTROL_ALLOW_HEADERS,
"content-type, authorization, last-event-id".to_string(),
),
(header::ACCESS_CONTROL_MAX_AGE, "600".to_string()),
],
)
.into_response();
if wants_private_network {
resp.headers_mut().insert(
"access-control-allow-private-network",
axum::http::HeaderValue::from_static("true"),
);
}
resp
}
fn allow_origin(mut resp: Response, origin: Option<&str>) -> Response {
if let Some(o) = origin
&& let Ok(v) = axum::http::HeaderValue::from_str(o)
{
resp.headers_mut()
.insert(header::ACCESS_CONTROL_ALLOW_ORIGIN, v);
}
resp
}
async fn card(State(app): State<Arc<App>>) -> Response {
let bridge = Arc::clone(&app.bridge);
let v = tokio::task::spawn_blocking(move || {
bridge.call("GetAgentCard", json!({}), Principal::anonymous())
})
.await
.unwrap_or(Value::Null);
(
StatusCode::OK,
[(header::CONTENT_TYPE, "application/json")],
serde_json::to_vec(&v).unwrap_or_default(),
)
.into_response()
}
async fn rpc(
State(app): State<Arc<App>>,
axum::Extension(peer_id): axum::Extension<PeerId>,
axum::Extension(Peer(peer)): axum::Extension<Peer>,
headers: HeaderMap,
body: Bytes,
) -> Response {
let allowed = headers
.get(header::ORIGIN)
.and_then(|v| v.to_str().ok())
.map(str::to_string);
allow_origin(
dispatch(app, peer_id, peer, headers, body).await,
allowed.as_deref(),
)
}
async fn dispatch(
app: Arc<App>,
peer_id: PeerId,
peer: SocketAddr,
headers: HeaderMap,
body: Bytes,
) -> Response {
let origin = headers
.get(header::ORIGIN)
.and_then(|v| v.to_str().ok())
.map(str::to_string);
if let Some(o) = &origin
&& !origin_allowed(o, &app.extra_origins)
{
return (StatusCode::FORBIDDEN, "origin not allowed").into_response();
}
let Ok(req) = serde_json::from_slice::<Value>(&body) else {
return err(Value::Null, -32700, "invalid JSON");
};
let id = req.get("id").cloned().unwrap_or(Value::Null);
let method = req
.get("method")
.and_then(Value::as_str)
.unwrap_or("")
.to_string();
let params = req.get("params").cloned().unwrap_or_else(|| json!({}));
let bearer = headers
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|h| {
h.strip_prefix("Bearer ")
.or_else(|| h.strip_prefix("bearer "))
})
.map(str::to_string);
let Some(principal) = resolve(&app, &peer_id, peer, bearer.as_deref()) else {
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer")],
"",
)
.into_response();
};
let bare = method.strip_prefix("a2a.").unwrap_or(&method).to_string();
match bare.as_str() {
"GetAgentCard" | "agent/card" => {
return unary(&app, id, "GetAgentCard", json!({}), principal).await;
}
"GetExtendedAgentCard" | "agent/getAuthenticatedExtendedCard" => {
if principal.is_anonymous() {
return err(
id,
-32007,
"the extended card requires an authenticated caller",
);
}
return unary(&app, id, "GetExtendedAgentCard", json!({}), principal).await;
}
"Pair" | "interface.pair" => {
return unary(&app, id, "Pair", params, principal).await;
}
"SubscribeToEvents" => {
return match &app.bridge.feed() {
Some(feed) => {
if !principal.may("SubscribeToEvents", None) {
return err(id, -32003, "not authorized");
}
feed_stream(Arc::clone(feed), id, params, principal, app.stream_deadline)
}
None => err(
id,
-32004,
"the interface surface is disabled (set interface.enabled: true)",
),
};
}
_ => {}
}
if crate::a2a::principals::is_admin(&method) {
if !principal.is_operator() {
return err(id, -32003, "operator role required");
}
return unary(&app, id, &method, params, principal).await;
}
let op = params
.get("message")
.and_then(crate::runtime::a2a_server::command_op);
if !principal.may(&bare, op.as_deref()) {
app.log.warn(
"a2a.denied",
json!({"principal": principal.id, "method": bare, "op": op}),
);
return err(id, -32003, "not authorized");
}
if matches!(bare.as_str(), "SendMessage" | "SendStreamingMessage")
&& crate::runtime::a2a_server::command_op(¶ms["message"]).is_some()
{
let streaming = bare == "SendStreamingMessage";
return unary_maybe_streamed(&app, id, "SendMessage", params, principal, streaming).await;
}
let body = match bare.as_str() {
"SendMessage" | "SendStreamingMessage" => {
match normalize_send(&app.bridge, &req, ¶ms).await {
Some(rewritten) => Bytes::from(rewritten),
None => body,
}
}
_ => body,
};
let mut request = axum::http::Request::builder()
.method("POST")
.uri("/")
.body(Body::from(body))
.expect("build request");
*request.headers_mut() = headers;
request
.extensions_mut()
.insert(a2a_rs::port::AuthPrincipal::new(
principal.id.clone(),
"agentd".to_string(),
));
let protocol = app.protocol.clone();
ports::with_caller(principal, async move {
protocol
.oneshot(request)
.await
.unwrap_or_else(|_| err(Value::Null, -32603, "dispatch failed"))
})
.await
}
async fn normalize_send(bridge: &Arc<A2aBridge>, req: &Value, params: &Value) -> Option<Vec<u8>> {
if !params.is_object() || !params.get("message").is_some_and(Value::is_object) {
return None;
}
let mut req = req.clone();
let mut changed = false;
if params["message"]["taskId"]
.as_str()
.unwrap_or("")
.is_empty()
{
let bridge = Arc::clone(bridge);
if let Ok(v) = tokio::task::spawn_blocking(move || {
bridge.call("NewTaskId", json!({}), Principal::anonymous())
})
.await
&& let Some(id) = v.get("id").and_then(Value::as_str)
&& let Some(message) = param_object(&mut req, "message")
{
message.insert("taskId".to_string(), json!(id));
changed = true;
}
}
if let Some(blocking) = params["configuration"]["blocking"].as_bool()
&& params["configuration"]["returnImmediately"].is_null()
&& let Some(config) = param_object(&mut req, "configuration")
{
config.insert("returnImmediately".to_string(), json!(!blocking));
changed = true;
}
changed.then(|| serde_json::to_vec(&req).ok()).flatten()
}
fn param_object<'a>(
req: &'a mut Value,
field: &str,
) -> Option<&'a mut serde_json::Map<String, Value>> {
req.as_object_mut()?
.get_mut("params")?
.as_object_mut()?
.get_mut(field)?
.as_object_mut()
}
async fn unary(
app: &Arc<App>,
id: Value,
method: &str,
params: Value,
principal: Principal,
) -> Response {
unary_maybe_streamed(app, id, method, params, principal, false).await
}
async fn unary_maybe_streamed(
app: &Arc<App>,
id: Value,
method: &str,
params: Value,
principal: Principal,
streamed: bool,
) -> Response {
let bridge = Arc::clone(&app.bridge);
let method = method.to_string();
let v = tokio::task::spawn_blocking(move || bridge.call(&method, params, principal))
.await
.unwrap_or_else(|e| json!({"_error": {"code": -32603, "message": e.to_string()}}));
let envelope = match v.get("_error") {
Some(e) => json!({"jsonrpc": "2.0", "id": id, "error": e}),
None => json!({"jsonrpc": "2.0", "id": id, "result": v}),
};
if !streamed {
return json_response(envelope);
}
let frame = axum::response::sse::Event::default()
.id("1")
.data(serde_json::to_string(&envelope).unwrap_or_default());
axum::response::Sse::new(futures_util::stream::once(async move {
Ok::<_, std::convert::Infallible>(frame)
}))
.into_response()
}
fn json_response(v: Value) -> Response {
(
StatusCode::OK,
[(header::CONTENT_TYPE, "application/json")],
serde_json::to_vec(&v).unwrap_or_default(),
)
.into_response()
}
fn err(id: Value, code: i64, message: &str) -> Response {
json_response(json!({"jsonrpc": "2.0", "id": id, "error": {"code": code, "message": message}}))
}
fn resolve(
app: &Arc<App>,
peer_id: &PeerId,
peer: SocketAddr,
bearer: Option<&str>,
) -> Option<Principal> {
let a = &app.auth;
if let (Some(p), Some(b)) = (&a.pairing, bearer)
&& let Some(role) = p.check_bearer(b)
{
return Some(crate::runtime::a2a_server::paired_principal(role));
}
let loopback = peer.ip().is_loopback();
let mgmt = (!a.require_auth && loopback) || peer_id.presented || is_server_bearer(a, bearer);
if !a.require_auth {
return Some(app.bridge.principal_of(
true,
bearer,
peer_id.subject.clone(),
peer_id.sans.clone(),
));
}
if !mgmt && bearer.is_none() && a.pairing.is_none() {
return None;
}
Some(
app.bridge
.principal_of(mgmt, bearer, peer_id.subject.clone(), peer_id.sans.clone()),
)
}
fn is_server_bearer(a: &Auth, bearer: Option<&str>) -> bool {
match (&a.server_bearer, bearer) {
(Some(server), Some(got)) => crate::sha::ct_eq(server.as_bytes(), got.as_bytes()),
_ => false,
}
}
fn origin_allowed(origin: &str, extra: &[String]) -> bool {
if extra.iter().any(|o| o == origin || o == "*") {
return true;
}
let host = origin
.split("://")
.nth(1)
.unwrap_or(origin)
.split(':')
.next()
.unwrap_or("");
crate::net::http::is_loopback_host(host)
}
fn feed_stream(
feed: Arc<SharedFeed>,
id: Value,
params: Value,
principal: Principal,
deadline: Duration,
) -> Response {
let after = params
.get("fromSeq")
.or_else(|| params.get("after"))
.and_then(Value::as_u64)
.unwrap_or(0);
let (newest, oldest, dropped) = feed.bounds();
let evicted = after > 0 && dropped > 0 && after < oldest.saturating_sub(1);
let ahead = after > newest;
let resync = evicted || ahead;
let start = if resync { 0 } else { after };
let (tx, rx) = tokio::sync::mpsc::channel::<axum::response::sse::Event>(64);
let is_op = principal.is_operator();
let who = principal.id.clone();
tokio::spawn(async move {
let hello = json!({"hello": {
"seq": newest,
"resume": after,
"resync": resync,
"debug": feed.debug(),
"version": crate::VERSION,
}});
if tx.send(frame(&id, hello)).await.is_err() {
return;
}
let mut cursor = start;
let end = Instant::now() + deadline;
loop {
let (events, next) = feed.since(cursor, &who, is_op, 256);
cursor = next;
for ev in events {
if tx.send(frame(&id, json!({"event": ev}))).await.is_err() {
return; }
}
if Instant::now() >= end {
let bye = json!({"goodbye": {"seq": cursor, "reason": "deadline"}});
let _ = tx.send(frame(&id, bye)).await;
return;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
});
let stream =
tokio_stream::wrappers::ReceiverStream::new(rx).map(Ok::<_, std::convert::Infallible>);
axum::response::Sse::new(stream)
.keep_alive(axum::response::sse::KeepAlive::default())
.into_response()
}
fn frame(id: &Value, payload: Value) -> axum::response::sse::Event {
axum::response::sse::Event::default().data(
serde_json::to_string(&json!({"jsonrpc": "2.0", "id": id, "result": payload}))
.unwrap_or_default(),
)
}
#[cfg(test)]
mod tests {
use super::*;
fn stub_bridge() -> Arc<A2aBridge> {
let resolver =
crate::a2a::Resolver::build(&crate::config::v2::A2a::default(), &|_| None).unwrap();
let (tx, rx) = std::sync::mpsc::channel();
std::thread::spawn(move || {
while let Ok(crate::runtime::events::Event::A2a(req)) = rx.recv() {
let _ = req.reply.send(json!({"id": "task-stub"}));
}
});
A2aBridge::new(tx, resolver)
}
#[tokio::test]
async fn malformed_send_params_are_left_alone_rather_than_panicking() {
let bridge = stub_bridge();
for params in [
json!([]),
json!({"message": "hi"}),
json!({"message": 3}),
json!({"message": []}),
Value::Null,
json!("send"),
json!({}),
] {
let req =
json!({"jsonrpc": "2.0", "id": 1, "method": "message/send", "params": params});
let p = req.get("params").cloned().unwrap_or_else(|| json!({}));
assert_eq!(
normalize_send(&bridge, &req, &p).await,
None,
"params {p} must not be rewritten"
);
}
}
#[tokio::test]
async fn a_well_formed_send_is_still_normalised() {
let bridge = stub_bridge();
let req = json!({"jsonrpc": "2.0", "id": 1, "method": "message/send", "params": {
"message": {"messageId": "m1", "role": "user", "parts": [{"kind": "text", "text": "hi"}]},
"configuration": {"blocking": false},
}});
let params = req["params"].clone();
let out = normalize_send(&bridge, &req, ¶ms)
.await
.expect("a well-formed send is rewritten");
let v: Value = serde_json::from_slice(&out).unwrap();
assert_eq!(v["params"]["message"]["taskId"], json!("task-stub"));
assert_eq!(
v["params"]["configuration"]["returnImmediately"],
json!(true)
);
}
}