use bytes::Bytes;
use http_body_util::BodyExt;
use a2a_protocol_server::streaming::event_queue::new_in_memory_queue_with_capacity;
use a2a_protocol_server::streaming::{build_sse_response, EventQueueWriter};
use a2a_protocol_types::events::{StreamResponse, TaskStatusUpdateEvent};
use a2a_protocol_types::task::{ContextId, TaskId, TaskState, TaskStatus};
const TINY_CAPACITY: usize = 2;
fn status_event(n: usize) -> StreamResponse {
StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: TaskId::new(format!("task-{n}")),
context_id: ContextId::new("ctx-1"),
status: TaskStatus::new(TaskState::Working),
metadata: None,
})
}
async fn lagged_stream_body(jsonrpc_envelope: bool) -> String {
let (writer, reader) = new_in_memory_queue_with_capacity(TINY_CAPACITY);
for n in 0..(TINY_CAPACITY * 4) {
writer
.write(status_event(n))
.await
.expect("write should succeed");
}
drop(writer);
let envelope_id = if jsonrpc_envelope {
Some(Some(serde_json::json!(1)))
} else {
None
};
let response = build_sse_response(reader, None, None, envelope_id);
let bytes: Bytes = response
.into_body()
.collect()
.await
.expect("collect body")
.to_bytes();
String::from_utf8(bytes.to_vec()).expect("SSE body is UTF-8")
}
fn error_frame_data(body: &str) -> Option<serde_json::Value> {
let mut saw_error_event = false;
for line in body.lines() {
if line.trim() == "event: error" {
saw_error_event = true;
} else if saw_error_event {
if let Some(data) = line.strip_prefix("data: ") {
return serde_json::from_str(data).ok();
}
}
}
None
}
#[tokio::test]
async fn jsonrpc_stream_error_is_a_jsonrpc_error_response() {
let body = lagged_stream_body(true).await;
let data = error_frame_data(&body)
.unwrap_or_else(|| panic!("no `event: error` frame in body:\n{body}"));
assert_eq!(
data.get("jsonrpc").and_then(serde_json::Value::as_str),
Some("2.0"),
"error frame must carry the JSON-RPC version; body:\n{body}"
);
assert!(
data.get("error").is_some(),
"error frame must carry an `error` member; body:\n{body}"
);
assert!(
data.get("result").is_none(),
"§5 requires exactly one of result/error; body:\n{body}"
);
assert_eq!(
data.get("id"),
Some(&serde_json::json!(1)),
"§9.4.2: the envelope echoes the originating request id; body:\n{body}"
);
let err = &data["error"];
assert!(
err.get("data")
.and_then(|d| d.get("streamLagged"))
.is_some(),
"the streamLagged marker must survive enveloping; body:\n{body}"
);
}
#[tokio::test]
async fn jsonrpc_stream_error_deserializes_as_a_jsonrpc_response() {
use a2a_protocol_types::jsonrpc::JsonRpcResponse;
let body = lagged_stream_body(true).await;
let data = error_frame_data(&body).expect("error frame");
let raw = serde_json::to_string(&data).expect("re-serialize");
let parsed: JsonRpcResponse<serde_json::Value> =
serde_json::from_str(&raw).unwrap_or_else(|e| {
panic!("a conformant client must parse this frame, got: {e}\nframe: {raw}")
});
assert!(
matches!(parsed, JsonRpcResponse::Error(_)),
"frame should parse as a JSON-RPC error response, got: {parsed:?}"
);
}
#[tokio::test]
async fn rest_stream_error_stays_a_bare_a2a_error() {
let body = lagged_stream_body(false).await;
let data = error_frame_data(&body)
.unwrap_or_else(|| panic!("no `event: error` frame in body:\n{body}"));
assert!(
data.get("jsonrpc").is_none(),
"REST frames carry no JSON-RPC envelope; body:\n{body}"
);
assert!(
data.get("code").is_some() && data.get("message").is_some(),
"REST error frame should be a bare A2aError; body:\n{body}"
);
assert!(
data.get("data")
.and_then(|d| d.get("streamLagged"))
.is_some(),
"the streamLagged marker must be present; body:\n{body}"
);
}