use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use futures_util::stream::SplitSink;
use futures_util::{SinkExt, StreamExt};
use tokio::net::{TcpListener, TcpStream};
use tokio_tungstenite::tungstenite::handshake::server::{
ErrorResponse, Request as WsUpgradeRequest, Response as WsUpgradeResponse,
};
use tokio_tungstenite::tungstenite::Message as WsMessage;
use tokio_tungstenite::WebSocketStream;
use a2a_protocol_types::jsonrpc::{
JsonRpcError, JsonRpcErrorResponse, JsonRpcId, JsonRpcRequest, JsonRpcSuccessResponse,
JsonRpcVersion,
};
use crate::error::ServerError;
use crate::handler::{RequestHandler, SendMessageResult};
use crate::streaming::EventQueueReader;
const MAX_WS_MESSAGE_SIZE: usize = 4 * 1024 * 1024;
const DEFAULT_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
pub struct WebSocketDispatcher {
handler: Arc<RequestHandler>,
handshake_timeout: Duration,
require_version_header: bool,
}
impl WebSocketDispatcher {
#[must_use]
pub const fn new(handler: Arc<RequestHandler>) -> Self {
Self {
handler,
handshake_timeout: DEFAULT_HANDSHAKE_TIMEOUT,
require_version_header: true,
}
}
#[must_use]
pub const fn accept_missing_version_header(mut self) -> Self {
self.require_version_header = false;
self
}
#[must_use]
pub const fn with_handshake_timeout(mut self, timeout: Duration) -> Self {
self.handshake_timeout = timeout;
self
}
pub async fn serve(
self: Arc<Self>,
addr: impl tokio::net::ToSocketAddrs,
) -> std::io::Result<()> {
let listener = TcpListener::bind(addr).await?;
trace_info!(
addr = %listener.local_addr().unwrap_or_else(|_| SocketAddr::from(([0, 0, 0, 0], 0))),
"A2A WebSocket server listening"
);
self.accept_loop(listener).await;
Ok(())
}
pub async fn serve_with_addr(
self: Arc<Self>,
addr: impl tokio::net::ToSocketAddrs,
) -> std::io::Result<SocketAddr> {
let listener = TcpListener::bind(addr).await?;
let local_addr = listener.local_addr()?;
trace_info!(%local_addr, "A2A WebSocket server listening");
tokio::spawn(async move {
self.accept_loop(listener).await;
});
Ok(local_addr)
}
async fn accept_loop(self: Arc<Self>, listener: TcpListener) {
loop {
let (stream, _peer) = match listener.accept().await {
Ok(pair) => pair,
Err(e) => {
trace_warn!(error = %e, "accept() failed; retrying");
let backoff = crate::serve::accept_retry_backoff(&e);
tokio::time::sleep(backoff).await;
continue;
}
};
let dispatcher = Arc::clone(&self);
tokio::spawn(async move {
trace_debug!("WebSocket connection accepted");
if let Err(_e) = dispatcher.handle_connection(stream).await {
trace_warn!(error = %_e, "WebSocket connection error");
}
});
}
}
#[allow(clippy::result_large_err)]
async fn handle_connection(&self, stream: TcpStream) -> Result<(), WsError> {
let _ = stream.set_nodelay(true);
let ws_config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default()
.max_message_size(Some(MAX_WS_MESSAGE_SIZE))
.max_frame_size(Some(MAX_WS_MESSAGE_SIZE));
let mut upgrade_headers: Option<HashMap<String, String>> = None;
let require_version = self.require_version_header;
let callback = |req: &WsUpgradeRequest, resp: WsUpgradeResponse| {
check_a2a_version(req, require_version)?;
upgrade_headers = Some(extract_upgrade_headers(req));
Ok(resp)
};
let ws_stream = tokio::time::timeout(
self.handshake_timeout,
tokio_tungstenite::accept_hdr_async_with_config(stream, callback, Some(ws_config)),
)
.await
.map_err(|_| WsError::HandshakeTimeout)?
.map_err(WsError::Handshake)?;
let headers = Arc::new(upgrade_headers.unwrap_or_default());
let (writer, reader) = ws_stream.split();
let writer = Arc::new(tokio::sync::Mutex::new(writer));
self.read_loop(reader, &writer, &headers).await;
let mut w = writer.lock().await;
let _ = w.close().await;
drop(w);
Ok(())
}
async fn read_loop(
&self,
mut reader: futures_util::stream::SplitStream<WebSocketStream<TcpStream>>,
writer: &WsSink,
headers: &Arc<HashMap<String, String>>,
) {
let semaphore = Arc::new(tokio::sync::Semaphore::new(64));
while let Some(msg) = reader.next().await {
match msg {
Ok(WsMessage::Text(text)) => {
let Ok(permit) = semaphore.clone().try_acquire_owned() else {
let err_resp = JsonRpcErrorResponse::new(
best_effort_request_id(&text),
JsonRpcError::new(
-32000,
"server busy: too many concurrent requests".to_string(),
),
);
send_json(writer, &err_resp).await;
continue;
};
let writer = Arc::clone(writer);
let handler = Arc::clone(&self.handler);
let headers = Arc::clone(headers);
tokio::spawn(async move {
Box::pin(process_ws_message(&handler, &text, writer, &headers)).await;
drop(permit); });
}
Ok(WsMessage::Binary(_)) => {
let err_resp = JsonRpcErrorResponse::new(
None,
JsonRpcError::new(
-32700,
"binary frames are not supported; send JSON-RPC as text frames"
.to_string(),
),
);
send_json(writer, &err_resp).await;
}
Ok(WsMessage::Close(_)) | Err(_) => break,
Ok(_) => {}
}
}
}
}
fn extract_upgrade_headers(req: &WsUpgradeRequest) -> HashMap<String, String> {
let mut map: HashMap<String, String> = req
.headers()
.iter()
.filter_map(|(k, v)| {
v.to_str()
.ok()
.map(|val| (k.as_str().to_lowercase(), val.to_owned()))
})
.collect();
map.insert(":path".to_owned(), req.uri().path().to_owned());
map
}
#[allow(clippy::result_large_err)]
fn check_a2a_version(req: &WsUpgradeRequest, require: bool) -> Result<(), ErrorResponse> {
let value = req
.headers()
.get(a2a_protocol_types::A2A_VERSION_HEADER)
.and_then(|v| v.to_str().ok());
let v = value.unwrap_or("").trim();
if v.is_empty() {
if !require {
return Ok(());
}
} else {
let major = v.split('.').next().and_then(|s| s.parse::<u32>().ok());
if major == Some(1) {
return Ok(());
}
}
let a2a_err = a2a_protocol_types::error::A2aError::version_not_supported(if v.is_empty() {
"A2A version '0.3' is not supported by this server; expected '1.0' (send the A2A-Version header)"
.to_owned()
} else {
format!("unsupported A2A version: {v}; this server supports 1.x")
});
let mut error_obj = serde_json::json!({
"error": {
"code": a2a_err.code.http_status(),
"status": a2a_err.code.grpc_status(),
"message": a2a_err.message,
}
});
let details = a2a_err.error_info_data(None);
if !details.is_null() {
error_obj["error"]["details"] = details;
}
let body = error_obj.to_string();
let resp = tokio_tungstenite::tungstenite::http::Response::builder()
.status(400)
.header("content-type", "application/json")
.body(Some(body))
.unwrap_or_else(|_| {
let mut r = ErrorResponse::new(Some(String::new()));
*r.status_mut() = tokio_tungstenite::tungstenite::http::StatusCode::BAD_REQUEST;
r
});
Err(resp)
}
fn best_effort_request_id(text: &str) -> JsonRpcId {
let v: serde_json::Value = serde_json::from_str(text).ok()?;
match v.get("id") {
Some(serde_json::Value::Null) | None => None,
Some(id) => Some(id.clone()),
}
}
#[derive(Debug)]
enum WsError {
Handshake(tokio_tungstenite::tungstenite::Error),
HandshakeTimeout,
}
impl std::fmt::Display for WsError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Handshake(e) => write!(f, "WebSocket handshake failed: {e}"),
Self::HandshakeTimeout => write!(f, "WebSocket handshake timed out"),
}
}
}
type WsSink = Arc<tokio::sync::Mutex<SplitSink<WebSocketStream<TcpStream>, WsMessage>>>;
#[allow(clippy::too_many_lines)]
async fn process_ws_message(
handler: &RequestHandler,
text: &str,
writer: WsSink,
headers: &HashMap<String, String>,
) {
let rpc_req: JsonRpcRequest = match serde_json::from_str(text) {
Ok(req) => req,
Err(e) => {
let err_resp = JsonRpcErrorResponse::new(
None,
JsonRpcError::new(-32700, format!("parse error: {e}")),
);
send_json(&writer, &err_resp).await;
return;
}
};
let id = rpc_req.id.to_response_id();
match rpc_req.method.as_str() {
"SendMessage" => {
dispatch_send_message(handler, &rpc_req, false, headers, id, &writer).await;
}
"SendStreamingMessage" | "message/stream" => {
dispatch_send_message(handler, &rpc_req, true, headers, id, &writer).await;
}
"GetTask" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, p, hdr| {
Box::pin(async move {
let params: a2a_protocol_types::params::TaskQueryParams =
serde_json::from_value(p).map_err(|e| {
a2a_protocol_types::error::A2aError::invalid_params(e.to_string())
})?;
h.on_get_task(params, Some(hdr))
.await
.map(|r| serde_json::to_value(&r).unwrap_or_default())
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
"ListTasks" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, p, hdr| {
Box::pin(async move {
let params: a2a_protocol_types::params::ListTasksParams =
serde_json::from_value(p).map_err(|e| {
a2a_protocol_types::error::A2aError::invalid_params(e.to_string())
})?;
h.on_list_tasks(params, Some(hdr))
.await
.map(|r| serde_json::to_value(&r).unwrap_or_default())
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
"CancelTask" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, p, hdr| {
Box::pin(async move {
let params: a2a_protocol_types::params::CancelTaskParams =
serde_json::from_value(p).map_err(|e| {
a2a_protocol_types::error::A2aError::invalid_params(e.to_string())
})?;
h.on_cancel_task(params, Some(hdr))
.await
.map(|r| serde_json::to_value(&r).unwrap_or_default())
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
"SubscribeToTask" => {
let params = match parse_params::<a2a_protocol_types::params::TaskIdParams>(
rpc_req.params.as_ref(),
) {
Ok(p) => p,
Err(e) => {
send_error(&writer, id, &e).await;
return;
}
};
match handler.on_resubscribe(params, Some(headers)).await {
Ok(reader) => {
stream_events(&writer, reader, id).await;
}
Err(e) => {
send_error(&writer, id, &e).await;
}
}
}
"CreateTaskPushNotificationConfig" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, p, hdr| {
Box::pin(async move {
let params: a2a_protocol_types::push::TaskPushNotificationConfig =
serde_json::from_value(p).map_err(|e| {
a2a_protocol_types::error::A2aError::invalid_params(e.to_string())
})?;
h.on_set_push_config(params, Some(hdr))
.await
.map(|r| serde_json::to_value(&r).unwrap_or_default())
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
"GetTaskPushNotificationConfig" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, p, hdr| {
Box::pin(async move {
let params: a2a_protocol_types::params::GetPushConfigParams =
serde_json::from_value(p).map_err(|e| {
a2a_protocol_types::error::A2aError::invalid_params(e.to_string())
})?;
h.on_get_push_config(params, Some(hdr))
.await
.map(|r| serde_json::to_value(&r).unwrap_or_default())
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
"ListTaskPushNotificationConfigs" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, p, hdr| {
Box::pin(async move {
let params: a2a_protocol_types::params::ListPushConfigsParams =
serde_json::from_value(p).map_err(|e| {
a2a_protocol_types::error::A2aError::invalid_params(e.to_string())
})?;
h.on_list_push_configs(¶ms.task_id, params.tenant.as_deref(), Some(hdr))
.await
.map(|configs| {
let resp = a2a_protocol_types::responses::ListPushConfigsResponse {
configs,
next_page_token: None,
};
serde_json::to_value(&resp).unwrap_or_default()
})
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
"DeleteTaskPushNotificationConfig" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, p, hdr| {
Box::pin(async move {
let params: a2a_protocol_types::params::DeletePushConfigParams =
serde_json::from_value(p).map_err(|e| {
a2a_protocol_types::error::A2aError::invalid_params(e.to_string())
})?;
h.on_delete_push_config(params, Some(hdr))
.await
.map(|()| serde_json::json!({}))
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
"GetExtendedAgentCard" => {
dispatch_simple(handler, &rpc_req, id, headers, &writer, |h, _p, hdr| {
Box::pin(async move {
h.on_get_extended_agent_card(Some(hdr))
.await
.map(|r| serde_json::to_value(&r).unwrap_or_default())
.map_err(|e| e.to_a2a_error())
})
})
.await;
}
other => {
let err = ServerError::MethodNotFound(other.to_owned());
send_error(&writer, id, &err).await;
}
}
}
async fn dispatch_send_message(
handler: &RequestHandler,
rpc_req: &JsonRpcRequest,
streaming: bool,
headers: &HashMap<String, String>,
id: JsonRpcId,
writer: &WsSink,
) {
let params = match parse_params::<a2a_protocol_types::params::MessageSendParams>(
rpc_req.params.as_ref(),
) {
Ok(p) => p,
Err(e) => {
send_error(writer, id, &e).await;
return;
}
};
match handler
.on_send_message(params, streaming, Some(headers))
.await
{
Ok(SendMessageResult::Response(resp)) => {
let result = serde_json::to_value(&resp).unwrap_or(serde_json::Value::Null);
let success = JsonRpcSuccessResponse {
jsonrpc: JsonRpcVersion,
id,
result,
};
send_json(writer, &success).await;
}
Ok(SendMessageResult::Stream(reader)) => {
stream_events(writer, reader, id).await;
}
Err(e) => {
send_error(writer, id, &e).await;
}
}
}
async fn stream_events(
writer: &WsSink,
mut reader: crate::streaming::InMemoryQueueReader,
id: JsonRpcId,
) {
while let Some(event) = reader.read().await {
match event {
Ok(stream_resp) => {
let envelope = JsonRpcSuccessResponse {
jsonrpc: JsonRpcVersion,
id: id.clone(),
result: stream_resp,
};
let json = serde_json::to_string(&envelope).unwrap_or_default();
let mut w = writer.lock().await;
if w.send(WsMessage::Text(json.into())).await.is_err() {
return; }
drop(w);
}
Err(e) => {
let err_resp =
JsonRpcErrorResponse::new(id.clone(), JsonRpcError::new(-32000, e.to_string()));
send_json(writer, &err_resp).await;
return;
}
}
}
let success = JsonRpcSuccessResponse {
jsonrpc: JsonRpcVersion,
id,
result: serde_json::json!({"status": "stream_complete"}),
};
send_json(writer, &success).await;
}
async fn dispatch_simple<'a, F>(
handler: &'a RequestHandler,
rpc_req: &JsonRpcRequest,
id: JsonRpcId,
headers: &'a HashMap<String, String>,
writer: &WsSink,
f: F,
) where
F: FnOnce(
&'a RequestHandler,
serde_json::Value,
&'a HashMap<String, String>,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<serde_json::Value, a2a_protocol_types::error::A2aError>,
> + Send
+ 'a,
>,
>,
{
let params = rpc_req.params.clone().unwrap_or(serde_json::Value::Null);
match f(handler, params, headers).await {
Ok(result) => {
let success = JsonRpcSuccessResponse {
jsonrpc: JsonRpcVersion,
id,
result,
};
send_json(writer, &success).await;
}
Err(e) => {
let err_resp =
JsonRpcErrorResponse::new(id, JsonRpcError::new(e.code.as_i32(), e.message));
send_json(writer, &err_resp).await;
}
}
}
async fn send_json<T: serde::Serialize + Sync>(writer: &WsSink, value: &T) {
let json = serde_json::to_string(value).unwrap_or_default();
let mut w = writer.lock().await;
let _ = w.send(WsMessage::Text(json.into())).await;
drop(w);
}
async fn send_error(writer: &WsSink, id: JsonRpcId, err: &ServerError) {
let a2a_err = err.to_a2a_error();
let resp = JsonRpcErrorResponse::new(
id,
JsonRpcError::new(a2a_err.code.as_i32(), a2a_err.message),
);
send_json(writer, &resp).await;
}
fn parse_params<T: serde::de::DeserializeOwned>(
params: Option<&serde_json::Value>,
) -> Result<T, ServerError> {
let value = params.cloned().unwrap_or(serde_json::Value::Null);
serde_json::from_value(value)
.map_err(|e| ServerError::InvalidParams(format!("invalid params: {e}")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_params_with_valid_json() {
let value = Some(serde_json::json!({"id": "task-1"}));
let result: Result<a2a_protocol_types::params::TaskQueryParams, _> =
parse_params(value.as_ref());
assert!(result.is_ok());
assert_eq!(result.unwrap().id, "task-1");
}
#[test]
fn parse_params_with_none_returns_error() {
let result: Result<a2a_protocol_types::params::TaskQueryParams, _> = parse_params(None);
assert!(result.is_err());
}
#[test]
fn parse_params_with_wrong_type_returns_error() {
let value = Some(serde_json::json!("not an object"));
let result: Result<a2a_protocol_types::params::TaskQueryParams, _> =
parse_params(value.as_ref());
assert!(result.is_err());
}
#[test]
fn ws_error_display_contains_message() {
let err = WsError::Handshake(tokio_tungstenite::tungstenite::Error::ConnectionClosed);
let s = err.to_string();
assert!(s.contains("WebSocket handshake failed"));
}
#[test]
fn ws_error_display_handshake_timeout() {
let s = WsError::HandshakeTimeout.to_string();
assert!(s.contains("timed out"), "got: {s}");
}
#[test]
fn best_effort_request_id_extracts_string_and_number() {
assert_eq!(
best_effort_request_id(r#"{"jsonrpc":"2.0","id":"req-1","method":"GetTask"}"#),
Some(serde_json::json!("req-1"))
);
assert_eq!(
best_effort_request_id(r#"{"jsonrpc":"2.0","id":7,"method":"GetTask"}"#),
Some(serde_json::json!(7))
);
}
#[test]
fn best_effort_request_id_none_for_missing_null_or_invalid() {
assert_eq!(best_effort_request_id(r#"{"jsonrpc":"2.0"}"#), None);
assert_eq!(best_effort_request_id(r#"{"id":null}"#), None);
assert_eq!(best_effort_request_id("not json {{"), None);
}
#[test]
fn websocket_dispatcher_new() {
use crate::agent_executor;
use crate::RequestHandlerBuilder;
use std::sync::Arc;
struct DummyExec;
agent_executor!(DummyExec, |_ctx, _queue| async { Ok(()) });
let handler = Arc::new(RequestHandlerBuilder::new(DummyExec).build().unwrap());
let _dispatcher = WebSocketDispatcher::new(handler);
}
use crate::agent_executor;
use crate::RequestHandlerBuilder;
use a2a_protocol_types::events::{StreamResponse, TaskStatusUpdateEvent};
use a2a_protocol_types::task::{ContextId, TaskState, TaskStatus};
use futures_util::{SinkExt, StreamExt};
struct EchoExec;
agent_executor!(EchoExec, |ctx, queue| async {
queue
.write(StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ContextId::new(ctx.context_id.clone()),
status: TaskStatus::new(TaskState::Working),
metadata: None,
}))
.await?;
queue
.write(StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ContextId::new(ctx.context_id.clone()),
status: TaskStatus::new(TaskState::Completed),
metadata: None,
}))
.await?;
Ok(())
});
async fn spawn_ws_server() -> std::net::SocketAddr {
let handler = Arc::new(RequestHandlerBuilder::new(EchoExec).build().unwrap());
let dispatcher = Arc::new(WebSocketDispatcher::new(handler));
dispatcher
.serve_with_addr("127.0.0.1:0")
.await
.expect("bind to port 0")
}
async fn ws_connect(
addr: std::net::SocketAddr,
) -> tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>
{
use tokio_tungstenite::tungstenite::client::IntoClientRequest as _;
let mut req = format!("ws://{addr}").into_client_request().expect("url");
req.headers_mut()
.insert("a2a-version", "1.0".parse().expect("header"));
let (ws, _) = tokio_tungstenite::connect_async(req)
.await
.expect("ws connect");
ws
}
async fn read_text(
ws: &mut tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
) -> String {
let msg = tokio::time::timeout(std::time::Duration::from_secs(5), ws.next())
.await
.expect("timeout waiting for WS frame")
.expect("stream ended")
.expect("ws error");
msg.into_text()
.expect("not a text frame")
.as_str()
.to_owned()
}
fn send_message_json(id: &str) -> String {
serde_json::json!({
"jsonrpc": "2.0",
"method": "SendMessage",
"id": id,
"params": {
"message": {
"messageId": "msg-1",
"role": "ROLE_USER",
"parts": [{"text": "hello"}]
}
}
})
.to_string()
}
#[tokio::test]
async fn ws_send_message_success() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
ws.send(WsMessage::Text(send_message_json("sm-1").into()))
.await
.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert_eq!(v["id"], "sm-1");
assert!(v.get("result").is_some(), "expected result key: {text}");
}
#[tokio::test]
async fn ws_get_task_not_found() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "GetTask",
"id": "gt-1",
"params": {"id": "nonexistent"}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(v.get("error").is_some(), "expected error: {text}");
}
#[tokio::test]
async fn ws_list_tasks_success() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "ListTasks",
"id": "lt-1",
"params": {}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert_eq!(v["id"], "lt-1");
assert!(v.get("result").is_some(), "expected result: {text}");
}
#[tokio::test]
async fn ws_cancel_task_not_found() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "CancelTask",
"id": "ct-1",
"params": {"id": "nonexistent"}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(v.get("error").is_some(), "expected error: {text}");
}
#[tokio::test]
async fn ws_subscribe_task_not_found() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "SubscribeToTask",
"id": "sub-1",
"params": {"id": "nonexistent"}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(v.get("error").is_some(), "expected error: {text}");
}
#[tokio::test]
async fn ws_unknown_method_error() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "FooBar",
"id": "unk-1",
"params": {}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(v.get("error").is_some(), "expected error: {text}");
let msg = v["error"]["message"].as_str().unwrap_or("");
assert!(
msg.to_lowercase().contains("method")
|| msg.to_lowercase().contains("not found")
|| msg.to_lowercase().contains("unsupported"),
"error message should mention method not found: {msg}"
);
}
#[tokio::test]
async fn ws_invalid_json_parse_error() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
ws.send(WsMessage::Text("this is not json {{".into()))
.await
.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert_eq!(v["error"]["code"], -32700, "expected parse error code");
}
#[tokio::test]
async fn ws_oversized_message_rejected() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let big = "x".repeat(4 * 1024 * 1024 + 1);
if ws.send(WsMessage::Text(big.into())).await.is_ok() {
let outcome = tokio::time::timeout(std::time::Duration::from_secs(5), ws.next())
.await
.expect("server should react to the oversized message");
match outcome {
None | Some(Err(_) | Ok(WsMessage::Close(_))) => {}
Some(Ok(frame)) => panic!(
"server must not answer an oversized message with a frame, got: {frame:?}"
),
}
}
}
#[tokio::test]
async fn ws_large_message_under_cap_still_processed() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let big = "x".repeat(3 * 1024 * 1024);
ws.send(WsMessage::Text(big.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert_eq!(v["error"]["code"], -32700, "expected parse error: {text}");
}
#[tokio::test]
async fn ws_ping_pong_response() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
ws.send(WsMessage::Ping(vec![42, 43].into())).await.unwrap();
let pong = tokio::time::timeout(std::time::Duration::from_secs(3), async {
loop {
let msg = ws.next().await.unwrap().unwrap();
if let WsMessage::Pong(data) = msg {
return data;
}
}
})
.await
.expect("should get pong within 3s");
assert_eq!(pong, vec![42, 43]);
}
#[tokio::test]
async fn ws_get_task_invalid_params() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "GetTask",
"id": "gti-1",
"params": {"wrong_field": 123}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(
v.get("error").is_some(),
"expected error for bad params: {text}"
);
}
#[tokio::test]
async fn ws_send_streaming_message_events() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "SendStreamingMessage",
"id": "ssm-1",
"params": {
"message": {
"messageId": "msg-stream-1",
"role": "ROLE_USER",
"parts": [{"text": "stream me"}]
}
}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let mut frames = Vec::new();
let timeout = tokio::time::timeout(std::time::Duration::from_secs(5), async {
loop {
let msg = ws.next().await.unwrap().unwrap();
let text = msg.into_text().unwrap();
let done = text.contains("stream_complete");
frames.push(text);
if done {
break;
}
}
});
timeout.await.expect("streaming should complete within 5s");
assert!(
frames.len() >= 3,
"expected >= 3 frames, got {}: {:?}",
frames.len(),
frames
);
assert!(frames.last().unwrap().contains("stream_complete"));
}
#[tokio::test]
async fn ws_send_message_invalid_params() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "SendMessage",
"id": "smi-1",
"params": {"not_message": true}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(
v.get("error").is_some(),
"expected error for bad send params: {text}"
);
}
#[tokio::test]
async fn ws_subscribe_invalid_params() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "SubscribeToTask",
"id": "subi-1",
"params": {}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(
v.get("error").is_some(),
"expected error for bad subscribe params: {text}"
);
}
#[tokio::test]
async fn ws_cancel_task_invalid_params() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "CancelTask",
"id": "cti-1",
"params": {"wrong": 1}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert!(v.get("error").is_some(), "expected error: {text}");
}
#[tokio::test]
async fn ws_list_tasks_with_filters() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "ListTasks",
"id": "ltf-1",
"params": {
"contextId": "ctx-1",
"pageSize": 10
}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert_eq!(v["id"], "ltf-1");
assert!(v.get("result").is_some(), "expected result: {text}");
}
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
async fn ws_call(
ws: &mut tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
req: serde_json::Value,
) -> serde_json::Value {
ws.send(WsMessage::Text(req.to_string().into()))
.await
.expect("send");
let text = read_text(ws).await;
serde_json::from_str(&text).expect("response should be JSON")
}
#[tokio::test]
async fn ws_legacy_method_names_rejected() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
for legacy in ["message/send", "tasks/list", "tasks/get"] {
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": legacy,
"id": format!("legacy-{legacy}"),
"params": {}
}),
)
.await;
assert_eq!(
v["error"]["code"].as_i64(),
Some(-32601),
"v0.3-style name {legacy} must be MethodNotFound: {v}"
);
}
}
#[tokio::test]
#[allow(clippy::too_many_lines)]
async fn ws_push_config_methods_routed() {
use crate::push::PushSender;
use a2a_protocol_types::push::TaskPushNotificationConfig;
struct NoopSender;
impl PushSender for NoopSender {
fn send<'a>(
&'a self,
_url: &'a str,
_event: &'a StreamResponse,
_config: &'a TaskPushNotificationConfig,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = a2a_protocol_types::error::A2aResult<()>>
+ Send
+ 'a,
>,
> {
Box::pin(async { Ok(()) })
}
fn allows_private_urls(&self) -> bool {
true
}
}
let handler = Arc::new(
RequestHandlerBuilder::new(EchoExec)
.with_push_sender(NoopSender)
.build()
.unwrap(),
);
let dispatcher = Arc::new(WebSocketDispatcher::new(handler));
let addr = dispatcher
.serve_with_addr("127.0.0.1:0")
.await
.expect("bind");
let mut ws = ws_connect(addr).await;
let v = ws_call(
&mut ws,
serde_json::from_str::<serde_json::Value>(&send_message_json("pc-0")).unwrap(),
)
.await;
let task_id = v["result"]["task"]["id"]
.as_str()
.expect("task id in send result")
.to_owned();
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": "CreateTaskPushNotificationConfig",
"id": "pc-1",
"params": {
"taskId": task_id,
"url": "https://example.com/hook"
}
}),
)
.await;
assert!(v.get("result").is_some(), "set push config failed: {v}");
let config_id = v["result"]["id"]
.as_str()
.expect("server-assigned config id")
.to_owned();
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": "GetTaskPushNotificationConfig",
"id": "pc-2",
"params": {"taskId": task_id, "id": config_id}
}),
)
.await;
assert!(v.get("result").is_some(), "get push config failed: {v}");
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": "ListTaskPushNotificationConfigs",
"id": "pc-3",
"params": {"taskId": task_id}
}),
)
.await;
assert!(v.get("result").is_some(), "list push configs failed: {v}");
assert!(
v["result"]["configs"].is_array(),
"expected configs array: {v}"
);
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": "DeleteTaskPushNotificationConfig",
"id": "pc-4",
"params": {"taskId": task_id, "id": config_id}
}),
)
.await;
assert!(v.get("result").is_some(), "delete push config failed: {v}");
}
#[tokio::test]
async fn ws_get_extended_agent_card_routed() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": "GetExtendedAgentCard",
"id": "card-1",
"params": {}
}),
)
.await;
let err = v.get("error").expect("expected an error response");
assert_ne!(
err["code"], -32601,
"GetExtendedAgentCard must be routed, got: {v}"
);
}
#[tokio::test]
async fn ws_upgrade_headers_drive_tenant_resolution() {
use crate::tenant_resolver::HeaderTenantResolver;
let handler = Arc::new(
RequestHandlerBuilder::new(EchoExec)
.with_tenant_resolver(HeaderTenantResolver::default())
.require_resolved_tenant()
.build()
.unwrap(),
);
let dispatcher = Arc::new(WebSocketDispatcher::new(handler));
let addr = dispatcher
.serve_with_addr("127.0.0.1:0")
.await
.expect("bind");
let mut ws = ws_connect(addr).await;
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": "ListTasks",
"id": "t-1",
"params": {}
}),
)
.await;
let err = v.get("error").expect("headerless request must be rejected");
let msg = err["message"].as_str().unwrap_or("");
assert!(
msg.contains("tenant"),
"expected strict-tenancy rejection, got: {v}"
);
let mut req = format!("ws://{addr}").into_client_request().unwrap();
req.headers_mut()
.insert("a2a-version", "1.0".parse().unwrap());
req.headers_mut()
.insert("x-tenant-id", "acme".parse().unwrap());
let (mut ws, _) = tokio_tungstenite::connect_async(req)
.await
.expect("connect");
let v = ws_call(
&mut ws,
serde_json::json!({
"jsonrpc": "2.0",
"method": "ListTasks",
"id": "t-2",
"params": {}
}),
)
.await;
assert!(
v.get("result").is_some(),
"tenant header on the upgrade request must reach the resolver: {v}"
);
}
#[tokio::test]
async fn ws_version_mismatch_rejects_handshake() {
let addr = spawn_ws_server().await;
let mut req = format!("ws://{addr}").into_client_request().unwrap();
req.headers_mut()
.insert("a2a-version", "2.0".parse().unwrap());
let outcome = tokio_tungstenite::connect_async(req).await;
assert!(
outcome.is_err(),
"handshake with A2A-Version 2.0 must be rejected"
);
let mut req = format!("ws://{addr}").into_client_request().unwrap();
req.headers_mut()
.insert("a2a-version", "1.0".parse().unwrap());
assert!(
tokio_tungstenite::connect_async(req).await.is_ok(),
"handshake with A2A-Version 1.0 must succeed"
);
}
#[tokio::test]
async fn ws_missing_version_header_rejected_by_default() {
use tokio_tungstenite::tungstenite::client::IntoClientRequest as _;
let addr = spawn_ws_server().await;
let req = format!("ws://{addr}").into_client_request().unwrap();
assert!(
tokio_tungstenite::connect_async(req).await.is_err(),
"a handshake with no A2A-Version header must be rejected by default \
(§3.6.2 reads it as 0.3)"
);
}
#[tokio::test]
async fn ws_missing_version_header_accepted_with_optout() {
use tokio_tungstenite::tungstenite::client::IntoClientRequest as _;
let handler = Arc::new(RequestHandlerBuilder::new(EchoExec).build().unwrap());
let dispatcher =
Arc::new(WebSocketDispatcher::new(handler).accept_missing_version_header());
let addr = dispatcher
.serve_with_addr("127.0.0.1:0")
.await
.expect("bind to port 0");
let req = format!("ws://{addr}").into_client_request().unwrap();
assert!(
tokio_tungstenite::connect_async(req).await.is_ok(),
"accept_missing_version_header() must restore the tolerant behaviour"
);
}
#[tokio::test]
async fn ws_version_rejection_body_carries_error_details() {
use tokio_tungstenite::tungstenite::client::IntoClientRequest as _;
use tokio_tungstenite::tungstenite::Error as WsError;
let addr = spawn_ws_server().await;
let mut req = format!("ws://{addr}").into_client_request().unwrap();
req.headers_mut()
.insert("a2a-version", "2.0".parse().unwrap());
match tokio_tungstenite::connect_async(req).await {
Err(WsError::Http(resp)) => {
assert_eq!(resp.status(), 400, "version rejection is HTTP 400");
let body = resp.body().as_ref().expect("rejection carries a body");
let text = String::from_utf8_lossy(body);
let json: serde_json::Value =
serde_json::from_str(&text).expect("rejection body is JSON");
assert!(
!json["error"]["details"].is_null(),
"the AIP-193 details block must be present, got: {text}"
);
}
other => panic!("expected an HTTP 400 rejection, got: {other:?}"),
}
}
#[tokio::test]
async fn ws_lagged_stream_reports_server_error_code() {
struct FloodExec;
agent_executor!(FloodExec, |ctx, queue| async {
for _ in 0..512 {
queue
.write(StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ContextId::new(ctx.context_id.clone()),
status: TaskStatus::new(TaskState::Working),
metadata: None,
}))
.await?;
}
Ok(())
});
let handler = Arc::new(
RequestHandlerBuilder::new(FloodExec)
.with_event_queue_capacity(1)
.build()
.unwrap(),
);
let addr = Arc::new(WebSocketDispatcher::new(handler))
.serve_with_addr("127.0.0.1:0")
.await
.expect("bind to port 0");
let mut ws = ws_connect(addr).await;
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "SendStreamingMessage",
"id": "lag-1",
"params": {
"message": {
"messageId": "msg-lag-1",
"role": "ROLE_USER",
"parts": [{"text": "flood"}]
}
}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
let found = tokio::time::timeout(std::time::Duration::from_secs(10), async {
while let Some(Ok(msg)) = ws.next().await {
let Ok(text) = msg.into_text() else { continue };
let Ok(v) = serde_json::from_str::<serde_json::Value>(&text) else {
continue;
};
if let Some(code) = v["error"]["code"].as_i64() {
return Some(code);
}
}
None
})
.await
.expect("the lagged stream must produce an error frame within 10s");
assert_eq!(
found,
Some(-32000),
"a lagged stream must be reported with the JSON-RPC server-error \
code -32000, not a positive code"
);
}
#[tokio::test]
async fn ws_over_concurrency_limit_is_rejected_with_server_error_code() {
struct BlockingExec;
agent_executor!(BlockingExec, |_ctx, _queue| async {
std::future::pending::<()>().await;
Ok(())
});
let handler = Arc::new(RequestHandlerBuilder::new(BlockingExec).build().unwrap());
let addr = Arc::new(WebSocketDispatcher::new(handler))
.serve_with_addr("127.0.0.1:0")
.await
.expect("bind to port 0");
let mut ws = ws_connect(addr).await;
for i in 0..65 {
let req = serde_json::json!({
"jsonrpc": "2.0",
"method": "SendMessage",
"id": format!("busy-{i}"),
"params": {
"message": {
"messageId": format!("msg-busy-{i}"),
"role": "ROLE_USER",
"parts": [{"text": "block"}]
}
}
})
.to_string();
ws.send(WsMessage::Text(req.into())).await.unwrap();
}
let code = tokio::time::timeout(std::time::Duration::from_secs(10), async {
while let Some(Ok(msg)) = ws.next().await {
let Ok(text) = msg.into_text() else { continue };
let Ok(v) = serde_json::from_str::<serde_json::Value>(&text) else {
continue;
};
if let Some(c) = v["error"]["code"].as_i64() {
assert!(
v["error"]["message"]
.as_str()
.is_some_and(|m| m.contains("server busy")),
"expected the back-pressure rejection, got: {v}"
);
return Some(c);
}
}
None
})
.await
.expect("the over-limit request must be answered within 10s");
assert_eq!(
code,
Some(-32000),
"back-pressure must be reported with the JSON-RPC server-error \
code -32000, not a positive code"
);
}
#[tokio::test]
async fn ws_handshake_timeout_disconnects_stalled_peer() {
let handler = Arc::new(RequestHandlerBuilder::new(EchoExec).build().unwrap());
let dispatcher = Arc::new(
WebSocketDispatcher::new(handler)
.with_handshake_timeout(std::time::Duration::from_millis(200)),
);
let addr = dispatcher
.serve_with_addr("127.0.0.1:0")
.await
.expect("bind");
let mut stream = tokio::net::TcpStream::connect(addr).await.expect("tcp");
let mut buf = [0u8; 16];
let read = tokio::time::timeout(
std::time::Duration::from_secs(5),
tokio::io::AsyncReadExt::read(&mut stream, &mut buf),
)
.await
.expect("server should close the stalled connection");
assert!(
matches!(read, Ok(0) | Err(_)),
"expected EOF/reset from server, got: {read:?}"
);
}
#[tokio::test]
async fn ws_binary_frame_gets_error_response() {
let addr = spawn_ws_server().await;
let mut ws = ws_connect(addr).await;
ws.send(WsMessage::Binary(vec![1, 2, 3].into()))
.await
.unwrap();
let text = read_text(&mut ws).await;
let v: serde_json::Value = serde_json::from_str(&text).unwrap();
assert_eq!(v["error"]["code"], -32700, "expected parse-error code: {v}");
assert!(
v["error"]["message"]
.as_str()
.unwrap_or("")
.contains("binary"),
"error should explain binary frames are unsupported: {v}"
);
}
}