use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use futures_util::{SinkExt, StreamExt};
use tokio::sync::{mpsc, oneshot, Mutex};
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::Message as WsMessage;
use uuid::Uuid;
use a2a_protocol_types::{JsonRpcRequest, JsonRpcResponse};
use crate::error::{ClientError, ClientResult};
use crate::streaming::EventStream;
use crate::transport::Transport;
enum PendingRequest {
Unary(oneshot::Sender<Result<String, ClientError>>),
Streaming(mpsc::Sender<crate::streaming::event_stream::BodyChunk>),
}
struct WriteCommand {
text: String,
request_id: String,
pending: PendingRequest,
}
#[derive(Debug, Clone)]
pub struct WebSocketTransportConfig {
pub request_timeout: Duration,
pub extra_headers: HashMap<String, String>,
pub max_message_size: usize,
}
impl Default for WebSocketTransportConfig {
fn default() -> Self {
Self {
request_timeout: Duration::from_secs(30),
extra_headers: HashMap::new(),
max_message_size: crate::transport::DEFAULT_MAX_RESPONSE_SIZE,
}
}
}
impl WebSocketTransportConfig {
#[must_use]
pub const fn with_request_timeout(mut self, timeout: Duration) -> Self {
self.request_timeout = timeout;
self
}
#[must_use]
pub fn with_extra_headers(mut self, headers: HashMap<String, String>) -> Self {
self.extra_headers = headers;
self
}
#[must_use]
pub const fn with_max_message_size(mut self, max_bytes: usize) -> Self {
self.max_message_size = max_bytes;
self
}
}
pub struct WebSocketTransport {
inner: Arc<Inner>,
}
struct Inner {
write_tx: mpsc::Sender<WriteCommand>,
pending: Arc<Mutex<HashMap<String, PendingRequest>>>,
closed: Arc<AtomicBool>,
endpoint: String,
request_timeout: Duration,
reader_handle: tokio::task::JoinHandle<()>,
writer_handle: tokio::task::JoinHandle<()>,
}
impl Drop for Inner {
fn drop(&mut self) {
self.reader_handle.abort();
self.writer_handle.abort();
}
}
impl WebSocketTransport {
pub async fn connect(endpoint: impl Into<String>) -> ClientResult<Self> {
Self::connect_with_options(endpoint, Duration::from_secs(30), &HashMap::new()).await
}
pub async fn connect_with_timeout(
endpoint: impl Into<String>,
request_timeout: Duration,
) -> ClientResult<Self> {
Self::connect_with_options(endpoint, request_timeout, &HashMap::new()).await
}
pub async fn connect_with_options(
endpoint: impl Into<String>,
request_timeout: Duration,
extra_headers: &HashMap<String, String>,
) -> ClientResult<Self> {
Self::connect_with_config(
endpoint,
WebSocketTransportConfig::default()
.with_request_timeout(request_timeout)
.with_extra_headers(extra_headers.clone()),
)
.await
}
#[allow(clippy::too_many_lines)]
pub async fn connect_with_config(
endpoint: impl Into<String>,
config: WebSocketTransportConfig,
) -> ClientResult<Self> {
let endpoint = endpoint.into();
validate_ws_url(&endpoint)?;
let mut ws_request = endpoint
.as_str()
.into_client_request()
.map_err(|e| ClientError::Transport(format!("WebSocket request build failed: {e}")))?;
ws_request.headers_mut().insert(
a2a_protocol_types::A2A_VERSION_HEADER,
tokio_tungstenite::tungstenite::http::HeaderValue::from_static(
a2a_protocol_types::A2A_VERSION,
),
);
for (k, v) in &config.extra_headers {
let name = k
.parse::<tokio_tungstenite::tungstenite::http::HeaderName>()
.map_err(|e| {
ClientError::Transport(format!("invalid WebSocket header name {k:?}: {e}"))
})?;
let val = v
.parse::<tokio_tungstenite::tungstenite::http::HeaderValue>()
.map_err(|_| {
ClientError::Transport(format!("invalid WebSocket header value for {k:?}"))
})?;
ws_request.headers_mut().insert(name, val);
}
let ws_config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default()
.max_message_size(Some(config.max_message_size))
.max_frame_size(Some(config.max_message_size));
let (ws_stream, _resp) =
tokio_tungstenite::connect_async_with_config(ws_request, Some(ws_config), true)
.await
.map_err(|e| ClientError::Transport(format!("WebSocket connect failed: {e}")))?;
let (ws_writer, ws_reader) = ws_stream.split();
let pending: Arc<Mutex<HashMap<String, PendingRequest>>> =
Arc::new(Mutex::new(HashMap::new()));
let closed = Arc::new(AtomicBool::new(false));
let (write_tx, mut write_rx) = mpsc::channel::<WriteCommand>(64);
let pending_for_writer = Arc::clone(&pending);
let closed_for_writer = Arc::clone(&closed);
let writer_handle = tokio::spawn(async move {
let mut ws_writer = ws_writer;
while let Some(cmd) = write_rx.recv().await {
{
let mut map = pending_for_writer.lock().await;
map.insert(cmd.request_id, cmd.pending);
}
if ws_writer
.send(WsMessage::Text(cmd.text.into()))
.await
.is_err()
{
fail_all_pending(&pending_for_writer, &closed_for_writer).await;
break;
}
}
});
let pending_for_reader = Arc::clone(&pending);
let closed_for_reader = Arc::clone(&closed);
let reader_handle = tokio::spawn(async move {
let mut ws_reader = ws_reader;
loop {
match ws_reader.next().await {
Some(Ok(WsMessage::Text(text))) => {
route_frame(&pending_for_reader, text.as_str()).await;
}
Some(Ok(WsMessage::Close(_)) | Err(_)) | None => break,
Some(Ok(_)) => {}
}
}
fail_all_pending(&pending_for_reader, &closed_for_reader).await;
});
Ok(Self {
inner: Arc::new(Inner {
write_tx,
pending,
closed,
endpoint,
request_timeout: config.request_timeout,
reader_handle,
writer_handle,
}),
})
}
#[must_use]
pub fn endpoint(&self) -> &str {
&self.inner.endpoint
}
async fn execute_request(
&self,
method: &str,
params: serde_json::Value,
extra_headers: &HashMap<String, String>,
) -> ClientResult<serde_json::Value> {
self.check_open()?;
warn_dropped_per_request_headers(method, extra_headers);
trace_info!(method, endpoint = %self.inner.endpoint, "sending WebSocket JSON-RPC request");
let rpc_req = build_rpc_request(method, params);
let request_id = rpc_req
.id
.as_value()
.and_then(|v| v.as_str())
.unwrap_or("")
.to_owned();
let body = serde_json::to_string(&rpc_req).map_err(ClientError::Serialization)?;
let (tx, rx) = oneshot::channel();
self.inner
.write_tx
.send(WriteCommand {
text: body,
request_id: request_id.clone(),
pending: PendingRequest::Unary(tx),
})
.await
.map_err(|_| ClientError::Transport("WebSocket writer task closed".into()))?;
let response_text = match tokio::time::timeout(self.inner.request_timeout, rx).await {
Ok(received) => received
.map_err(|_| ClientError::Transport("WebSocket reader task closed".into()))??,
Err(_elapsed) => {
self.inner.pending.lock().await.remove(&request_id);
return Err(ClientError::Timeout("WebSocket response timed out".into()));
}
};
let envelope: JsonRpcResponse<serde_json::Value> =
serde_json::from_str(&response_text).map_err(ClientError::Serialization)?;
match envelope {
JsonRpcResponse::Success(ok) => {
trace_info!(method, "WebSocket request succeeded");
Ok(ok.result)
}
JsonRpcResponse::Error(err) => {
trace_warn!(
method,
code = err.error.code,
"JSON-RPC error over WebSocket"
);
let a2a = crate::transport::map_jsonrpc_error(
err.error.code,
err.error.message,
err.error.data,
);
Err(ClientError::Protocol(a2a))
}
}
}
fn check_open(&self) -> ClientResult<()> {
if self.inner.closed.load(Ordering::Acquire) {
return Err(ClientError::Transport("WebSocket connection closed".into()));
}
Ok(())
}
async fn execute_streaming_request(
&self,
method: &str,
params: serde_json::Value,
extra_headers: &HashMap<String, String>,
) -> ClientResult<EventStream> {
self.check_open()?;
warn_dropped_per_request_headers(method, extra_headers);
trace_info!(method, endpoint = %self.inner.endpoint, "opening WebSocket stream");
let rpc_req = build_rpc_request(method, params);
let request_id = rpc_req
.id
.as_value()
.and_then(|v| v.as_str())
.unwrap_or("")
.to_owned();
let body = serde_json::to_string(&rpc_req).map_err(ClientError::Serialization)?;
let (tx, rx) = mpsc::channel::<crate::streaming::event_stream::BodyChunk>(64);
self.inner
.write_tx
.send(WriteCommand {
text: body,
request_id,
pending: PendingRequest::Streaming(tx),
})
.await
.map_err(|_| ClientError::Transport("WebSocket writer task closed".into()))?;
Ok(EventStream::new(rx).with_first_event_timeout(self.inner.request_timeout))
}
}
impl Transport for WebSocketTransport {
fn send_request<'a>(
&'a self,
method: &'a str,
params: serde_json::Value,
extra_headers: &'a HashMap<String, String>,
) -> Pin<Box<dyn Future<Output = ClientResult<serde_json::Value>> + Send + 'a>> {
Box::pin(self.execute_request(method, params, extra_headers))
}
fn send_streaming_request<'a>(
&'a self,
method: &'a str,
params: serde_json::Value,
extra_headers: &'a HashMap<String, String>,
) -> Pin<Box<dyn Future<Output = ClientResult<EventStream>> + Send + 'a>> {
Box::pin(self.execute_streaming_request(method, params, extra_headers))
}
}
impl std::fmt::Debug for WebSocketTransport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocketTransport")
.field("endpoint", &self.inner.endpoint)
.finish()
}
}
#[cfg_attr(not(feature = "tracing"), allow(unused_variables))]
fn warn_dropped_per_request_headers(method: &str, extra_headers: &HashMap<String, String>) {
if !extra_headers.is_empty() {
trace_warn!(
method,
header_count = extra_headers.len(),
"per-request headers are not sent over an established WebSocket connection; \
supply credentials at connect time via WebSocketTransport::connect_with_options"
);
}
}
async fn fail_all_pending(pending: &Mutex<HashMap<String, PendingRequest>>, closed: &AtomicBool) {
closed.store(true, Ordering::Release);
let entries: Vec<PendingRequest> = {
let mut map = pending.lock().await;
map.drain().map(|(_, v)| v).collect()
};
for entry in entries {
match entry {
PendingRequest::Unary(tx) => {
let _ = tx.send(Err(ClientError::Transport(
"WebSocket connection closed".into(),
)));
}
PendingRequest::Streaming(tx) => {
let _ = tx.try_send(Err(ClientError::Transport(
"WebSocket connection closed".into(),
)));
}
}
}
}
async fn route_frame(pending: &Arc<Mutex<HashMap<String, PendingRequest>>>, text: &str) {
let Some(request_id) = extract_jsonrpc_id(text) else {
return;
};
let streaming_tx = {
let mut map = pending.lock().await;
let tx = match map.get(&request_id) {
Some(PendingRequest::Unary(_)) => {
if let Some(PendingRequest::Unary(tx)) = map.remove(&request_id) {
let _ = tx.send(Ok(text.to_owned()));
}
return;
}
Some(PendingRequest::Streaming(tx)) => tx.clone(),
None => return,
};
drop(map);
tx
};
let sse_line = format!("data: {text}\n\n");
if streaming_tx
.send(Ok(hyper::body::Bytes::from(sse_line)))
.await
.is_err()
{
pending.lock().await.remove(&request_id);
return;
}
if is_stream_terminal(text) {
pending.lock().await.remove(&request_id);
}
}
fn extract_jsonrpc_id(text: &str) -> Option<String> {
let v: serde_json::Value = serde_json::from_str(text).ok()?;
match v.get("id") {
Some(serde_json::Value::String(s)) => Some(s.clone()),
Some(serde_json::Value::Number(n)) => Some(n.to_string()),
_ => None,
}
}
fn task_state_str_is_terminal(state: &str) -> bool {
serde_json::from_value::<a2a_protocol_types::TaskState>(serde_json::Value::String(
state.to_owned(),
))
.is_ok_and(a2a_protocol_types::TaskState::is_terminal)
}
fn is_stream_terminal(text: &str) -> bool {
let Ok(frame) = serde_json::from_str::<serde_json::Value>(text) else {
return false;
};
let has_terminal_state = |obj: &serde_json::Value| -> bool {
if let Some(status_update) = obj.get("statusUpdate") {
if let Some(status) = status_update.get("status") {
if let Some(state) = status.get("state").and_then(|s| s.as_str()) {
return task_state_str_is_terminal(state);
}
}
}
if let Some(status) = obj.get("status") {
if let Some(state) = status.get("state").and_then(|s| s.as_str()) {
return task_state_str_is_terminal(state);
}
}
false
};
if let Some(r) = frame.get("result") {
if r.get("stream_complete").is_some() {
return true;
}
if r.get("status").and_then(|s| s.as_str()) == Some("stream_complete") {
return true;
}
return has_terminal_state(r);
}
has_terminal_state(&frame)
}
fn build_rpc_request(method: &str, params: serde_json::Value) -> JsonRpcRequest {
let id = serde_json::Value::String(Uuid::new_v4().to_string());
JsonRpcRequest::with_params(id, method, params)
}
fn validate_ws_url(url: &str) -> ClientResult<()> {
if url.is_empty() {
return Err(ClientError::InvalidEndpoint("URL must not be empty".into()));
}
if !url.starts_with("ws://") && !url.starts_with("wss://") {
return Err(ClientError::InvalidEndpoint(format!(
"WebSocket URL must start with ws:// or wss://: {url}"
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validate_ws_url_rejects_empty() {
assert!(validate_ws_url("").is_err());
}
#[test]
fn with_extra_headers_sets_the_headers() {
let mut headers = HashMap::new();
headers.insert("authorization".to_string(), "Bearer tok".to_string());
headers.insert("x-custom".to_string(), "v".to_string());
let config = WebSocketTransportConfig::default().with_extra_headers(headers.clone());
assert_eq!(config.extra_headers, headers);
assert_eq!(
config
.extra_headers
.get("authorization")
.map(String::as_str),
Some("Bearer tok")
);
}
#[test]
fn validate_ws_url_rejects_http() {
assert!(validate_ws_url("http://localhost:8080").is_err());
}
#[test]
fn validate_ws_url_accepts_ws() {
assert!(validate_ws_url("ws://localhost:8080").is_ok());
}
#[test]
fn validate_ws_url_accepts_wss() {
assert!(validate_ws_url("wss://agent.example.com/a2a").is_ok());
}
#[test]
fn is_stream_terminal_completed_status() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"statusUpdate":{"status":{"state":"completed"}}}}"#;
assert!(is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_failed_status() {
let frame =
r#"{"jsonrpc":"2.0","id":"1","result":{"statusUpdate":{"status":{"state":"failed"}}}}"#;
assert!(is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_working_is_not_terminal() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"statusUpdate":{"status":{"state":"working"}}}}"#;
assert!(!is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_stream_complete_sentinel() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"stream_complete":true}}"#;
assert!(is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_artifact_not_terminal() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"artifactUpdate":{"artifact":{"id":"a1","parts":[]}}}}"#;
assert!(!is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_payload_containing_word_not_terminal() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"artifactUpdate":{"artifact":{"id":"a1","parts":[{"text":"task completed successfully"}]}}}}"#;
assert!(!is_stream_terminal(frame));
}
#[test]
fn build_rpc_request_has_method() {
let req = build_rpc_request("TestMethod", serde_json::json!({"key": "val"}));
assert_eq!(req.method, "TestMethod");
let params = req.params.expect("params should be present");
assert_eq!(params["key"], "val");
let id = req.id.as_value().expect("id should be present");
assert!(id.is_string(), "id should be a string UUID");
assert!(!id.as_str().unwrap().is_empty(), "id should not be empty");
}
#[test]
fn is_stream_terminal_invalid_json() {
assert!(!is_stream_terminal("not json"));
}
#[test]
fn is_stream_terminal_no_result() {
assert!(!is_stream_terminal(r#"{"jsonrpc":"2.0","id":"1"}"#));
}
#[test]
fn is_stream_terminal_task_level_completed() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"status":{"state":"completed"}}}"#;
assert!(is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_canceled() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"statusUpdate":{"status":{"state":"canceled"}}}}"#;
assert!(is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_rejected() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"statusUpdate":{"status":{"state":"rejected"}}}}"#;
assert!(is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_task_level_failed() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"status":{"state":"failed"}}}"#;
assert!(is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_non_string_state() {
let frame = r#"{"jsonrpc":"2.0","id":"1","result":{"status":{"state":42}}}"#;
assert!(!is_stream_terminal(frame));
}
#[test]
fn is_stream_terminal_canonical_screaming_snake_case() {
for state in [
"TASK_STATE_COMPLETED",
"TASK_STATE_FAILED",
"TASK_STATE_CANCELED",
"TASK_STATE_REJECTED",
] {
let frame = format!(
r#"{{"jsonrpc":"2.0","id":"1","result":{{"statusUpdate":{{"status":{{"state":"{state}"}}}}}}}}"#
);
assert!(
is_stream_terminal(&frame),
"canonical terminal state {state} not detected"
);
}
}
#[test]
fn is_stream_terminal_canonical_non_terminal() {
for state in ["TASK_STATE_WORKING", "TASK_STATE_SUBMITTED", "working"] {
let frame = format!(
r#"{{"jsonrpc":"2.0","id":"1","result":{{"status":{{"state":"{state}"}}}}}}"#
);
assert!(
!is_stream_terminal(&frame),
"non-terminal state {state} wrongly detected as terminal"
);
}
}
#[test]
fn validate_ws_url_rejects_https() {
assert!(validate_ws_url("https://example.com").is_err());
}
#[test]
fn validate_ws_url_error_message_contains_url() {
let err = validate_ws_url("http://bad").unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("http://bad") || msg.contains("ws://"));
}
#[test]
fn extract_jsonrpc_id_string() {
let id = extract_jsonrpc_id(r#"{"jsonrpc":"2.0","id":"abc","result":{}}"#);
assert_eq!(id.as_deref(), Some("abc"));
}
#[test]
fn extract_jsonrpc_id_number() {
let id = extract_jsonrpc_id(r#"{"jsonrpc":"2.0","id":42,"result":{}}"#);
assert_eq!(id.as_deref(), Some("42"));
}
#[test]
fn extract_jsonrpc_id_null_returns_none() {
let id = extract_jsonrpc_id(r#"{"jsonrpc":"2.0","id":null,"result":{}}"#);
assert!(id.is_none());
}
#[test]
fn extract_jsonrpc_id_missing_returns_none() {
let id = extract_jsonrpc_id(r#"{"jsonrpc":"2.0","result":{}}"#);
assert!(id.is_none());
}
#[tokio::test]
async fn timed_out_request_is_removed_from_pending_map() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
tokio::spawn(async move {
let Ok(mut ws) = tokio_tungstenite::accept_async(stream).await else {
return;
};
while let Some(Ok(_)) = ws.next().await {}
});
}
});
let transport = WebSocketTransport::connect_with_timeout(
format!("ws://{addr}"),
Duration::from_millis(100),
)
.await
.expect("connect");
let err = transport
.send_request("GetTask", serde_json::json!({"id": "t1"}), &HashMap::new())
.await
.expect_err("request must time out");
assert!(
matches!(err, ClientError::Timeout(_)),
"expected timeout, got: {err:?}"
);
assert!(
transport.inner.pending.lock().await.is_empty(),
"pending map must not retain timed-out requests"
);
}
async fn spawn_raw_ws_server<F, Fut>(per_conn: F) -> std::net::SocketAddr
where
F: Fn(tokio_tungstenite::WebSocketStream<tokio::net::TcpStream>) -> Fut
+ Send
+ Sync
+ 'static,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let per_conn = Arc::new(per_conn);
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
let per_conn = Arc::clone(&per_conn);
tokio::spawn(async move {
if let Ok(ws) = tokio_tungstenite::accept_async(stream).await {
per_conn(ws).await;
}
});
}
});
addr
}
#[tokio::test]
async fn dropping_transport_closes_connection() {
let (closed_tx, mut closed_rx) = mpsc::channel::<()>(1);
let closed_tx = Arc::new(closed_tx);
let addr = spawn_raw_ws_server(move |mut ws| {
let closed_tx = Arc::clone(&closed_tx);
async move {
while let Some(Ok(_)) = ws.next().await {}
let _ = closed_tx.send(()).await;
}
})
.await;
let transport = WebSocketTransport::connect(format!("ws://{addr}"))
.await
.expect("connect");
drop(transport);
tokio::time::timeout(Duration::from_secs(5), closed_rx.recv())
.await
.expect("server must observe the connection closing after drop")
.expect("channel open");
}
#[tokio::test]
async fn server_close_fails_pending_request_fast() {
let addr = spawn_raw_ws_server(|mut ws| async move {
let _ = ws.next().await;
let _ = ws.close(None).await;
})
.await;
let transport = WebSocketTransport::connect_with_timeout(
format!("ws://{addr}"),
Duration::from_secs(30),
)
.await
.expect("connect");
let start = std::time::Instant::now();
let err = transport
.send_request("GetTask", serde_json::json!({"id": "t1"}), &HashMap::new())
.await
.expect_err("request must fail when the server closes");
assert!(
matches!(err, ClientError::Transport(_)),
"expected transport error, got: {err:?}"
);
assert!(
start.elapsed() < Duration::from_secs(10),
"failure must be prompt, took {:?} against a 30s request timeout",
start.elapsed()
);
let err = transport
.send_request("GetTask", serde_json::json!({"id": "t2"}), &HashMap::new())
.await
.expect_err("dead transport must reject new requests");
assert!(
matches!(err, ClientError::Transport(_)),
"expected transport error, got: {err:?}"
);
}
#[tokio::test]
async fn oversized_incoming_frame_is_rejected() {
let addr = spawn_raw_ws_server(|mut ws| async move {
if let Some(Ok(_)) = ws.next().await {
let big = "x".repeat(64 * 1024);
let _ = ws
.send(tokio_tungstenite::tungstenite::Message::Text(big.into()))
.await;
}
while let Some(Ok(_)) = ws.next().await {}
})
.await;
let transport = WebSocketTransport::connect_with_config(
format!("ws://{addr}"),
WebSocketTransportConfig::default()
.with_request_timeout(Duration::from_secs(30))
.with_max_message_size(16 * 1024),
)
.await
.expect("connect");
let start = std::time::Instant::now();
let err = transport
.send_request("GetTask", serde_json::json!({"id": "t1"}), &HashMap::new())
.await
.expect_err("oversized frame must fail the request");
assert!(
matches!(err, ClientError::Transport(_)),
"expected transport error, got: {err:?}"
);
assert!(
start.elapsed() < Duration::from_secs(10),
"rejection must be prompt, took {:?}",
start.elapsed()
);
}
#[cfg(feature = "tracing")]
#[test]
fn warn_dropped_per_request_headers_warns_iff_headers_present() {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
struct CountingSubscriber(Arc<AtomicUsize>);
impl tracing::Subscriber for CountingSubscriber {
fn enabled(&self, _: &tracing::Metadata<'_>) -> bool {
true
}
fn new_span(&self, _: &tracing::span::Attributes<'_>) -> tracing::span::Id {
tracing::span::Id::from_u64(1)
}
fn record(&self, _: &tracing::span::Id, _: &tracing::span::Record<'_>) {}
fn record_follows_from(&self, _: &tracing::span::Id, _: &tracing::span::Id) {}
fn event(&self, _: &tracing::Event<'_>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
fn enter(&self, _: &tracing::span::Id) {}
fn exit(&self, _: &tracing::span::Id) {}
}
let count = Arc::new(AtomicUsize::new(0));
tracing::subscriber::with_default(CountingSubscriber(Arc::clone(&count)), || {
warn_dropped_per_request_headers("SendMessage", &HashMap::new());
assert_eq!(
count.load(Ordering::SeqCst),
0,
"must not warn when there are no per-request headers to drop"
);
let mut headers = HashMap::new();
headers.insert("authorization".to_owned(), "Bearer secret".to_owned());
warn_dropped_per_request_headers("SendMessage", &headers);
assert_eq!(
count.load(Ordering::SeqCst),
1,
"dropping a per-request header must emit a warning (never silent)"
);
});
}
}