use std::net::SocketAddr;
use std::pin::Pin;
use std::time::Duration;
use eventsource_stream::Eventsource;
use futures_core::Stream;
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use tokio_stream::StreamExt;
use zeph_common::net::resolve_and_validate;
use crate::error::A2aError;
use crate::ibct::{Ibct, IbctKey, ibct_scope_origin};
use crate::jsonrpc::{
JsonRpcRequest, JsonRpcResponse, METHOD_CANCEL_TASK, METHOD_GET_TASK, METHOD_SEND_MESSAGE,
METHOD_SEND_STREAMING_MESSAGE, SendMessageParams, TaskIdParams,
};
use crate::types::{Task, TaskArtifactUpdateEvent, TaskStatusUpdateEvent};
pub type TaskEventStream = Pin<Box<dyn Stream<Item = Result<TaskEvent, A2aError>> + Send>>;
#[non_exhaustive]
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum TaskEvent {
StatusUpdate(TaskStatusUpdateEvent),
ArtifactUpdate(TaskArtifactUpdateEvent),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SecurityPolicy {
pub require_tls: bool,
pub ssrf_protection: bool,
}
impl SecurityPolicy {
#[must_use]
pub const fn hardened() -> Self {
Self {
require_tls: true,
ssrf_protection: true,
}
}
#[must_use]
pub const fn permissive() -> Self {
Self {
require_tls: false,
ssrf_protection: false,
}
}
}
#[derive(Debug)]
struct PinnedTarget {
host: String,
addrs: Vec<SocketAddr>,
}
pub struct A2aClient {
client: reqwest::Client,
security: SecurityPolicy,
request_timeout: Duration,
ibct_key: Option<IbctKey>,
ibct_ttl: Duration,
}
impl A2aClient {
#[must_use]
pub fn new(client: reqwest::Client) -> Self {
Self {
client,
security: SecurityPolicy::permissive(),
request_timeout: Duration::from_secs(30),
ibct_key: None,
ibct_ttl: Duration::from_mins(5),
}
}
#[must_use]
pub fn with_security(mut self, policy: SecurityPolicy) -> Self {
self.security = policy;
self
}
#[must_use]
pub fn with_request_timeout(mut self, timeout: Duration) -> Self {
self.request_timeout = timeout;
self
}
#[must_use]
pub fn with_ibct_key(mut self, key: IbctKey) -> Self {
self.ibct_key = Some(key);
self
}
#[must_use]
pub fn with_ibct_ttl(mut self, ttl: Duration) -> Self {
self.ibct_ttl = ttl;
self
}
fn ibct_header_value(&self, endpoint: &str, task_id: &str) -> Option<String> {
let key = self.ibct_key.as_ref()?;
let scope = ibct_scope_origin(endpoint);
let token = match Ibct::issue(task_id, &scope, self.ibct_ttl, key) {
Ok(t) => t,
Err(e) => {
tracing::warn!("failed to issue IBCT token: {e}");
return None;
}
};
match token.encode() {
Ok(encoded) => Some(encoded),
Err(e) => {
tracing::warn!("failed to encode IBCT token: {e}");
None
}
}
}
#[tracing::instrument(name = "a2a.client.send_message", skip_all, err)]
pub async fn send_message(
&self,
endpoint: &str,
params: SendMessageParams,
token: Option<&str>,
) -> Result<Task, A2aError> {
let task_id = params.message.task_id.clone().unwrap_or_default();
self.rpc_call(endpoint, METHOD_SEND_MESSAGE, params, token, &task_id)
.await
}
#[tracing::instrument(name = "a2a.client.stream_message", skip_all, err)]
pub async fn stream_message(
&self,
endpoint: &str,
params: SendMessageParams,
token: Option<&str>,
) -> Result<TaskEventStream, A2aError> {
let task_id = params.message.task_id.clone().unwrap_or_default();
let pinned = self.validate_endpoint(endpoint).await?;
let request_client = self.request_client(pinned.as_ref())?;
let request = JsonRpcRequest::new(METHOD_SEND_STREAMING_MESSAGE, params);
let mut req = request_client.post(endpoint).json(&request);
if let Some(t) = token {
req = req.bearer_auth(t);
}
if let Some(ibct_header) = self.ibct_header_value(endpoint, &task_id) {
req = req.header("X-Zeph-IBCT", ibct_header);
}
let resp = tokio::time::timeout(self.request_timeout, req.send())
.await
.map_err(|_| A2aError::Timeout(self.request_timeout))?
.map_err(A2aError::Http)?;
if !resp.status().is_success() {
let status = resp.status();
let body = tokio::time::timeout(Duration::from_secs(5), resp.text())
.await
.unwrap_or(Ok(String::new()))
.unwrap_or_default();
let truncated = if body.len() > 256 {
format!("{}…", &body[..256])
} else {
body
};
return Err(A2aError::Stream(format!("HTTP {status}: {truncated}")));
}
let event_stream = resp.bytes_stream().eventsource();
let mapped = event_stream.filter_map(|event| match event {
Ok(event) => {
if event.data.is_empty() || event.data == "[DONE]" {
return None;
}
match serde_json::from_str::<JsonRpcResponse<TaskEvent>>(&event.data) {
Ok(rpc_resp) => match rpc_resp.into_result() {
Ok(task_event) => Some(Ok(task_event)),
Err(rpc_err) => Some(Err(A2aError::from(rpc_err))),
},
Err(e) => Some(Err(A2aError::Stream(format!(
"failed to parse SSE event: {e}"
)))),
}
}
Err(e) => Some(Err(A2aError::Stream(format!("SSE stream error: {e}")))),
});
Ok(Box::pin(mapped))
}
#[tracing::instrument(name = "a2a.client.get_task", skip_all, err)]
pub async fn get_task(
&self,
endpoint: &str,
params: TaskIdParams,
token: Option<&str>,
) -> Result<Task, A2aError> {
let task_id = params.id.clone();
self.rpc_call(endpoint, METHOD_GET_TASK, params, token, &task_id)
.await
}
#[tracing::instrument(name = "a2a.client.cancel_task", skip_all, err)]
pub async fn cancel_task(
&self,
endpoint: &str,
params: TaskIdParams,
token: Option<&str>,
) -> Result<Task, A2aError> {
let task_id = params.id.clone();
self.rpc_call(endpoint, METHOD_CANCEL_TASK, params, token, &task_id)
.await
}
#[tracing::instrument(name = "a2a.client.validate_endpoint", skip_all, err)]
async fn validate_endpoint(&self, endpoint: &str) -> Result<Option<PinnedTarget>, A2aError> {
if self.security.require_tls && !endpoint.starts_with("https://") {
return Err(A2aError::Security(format!(
"TLS required but endpoint uses HTTP: {endpoint}"
)));
}
if !self.security.ssrf_protection {
return Ok(None);
}
let url: url::Url = endpoint
.parse()
.map_err(|e| A2aError::Security(format!("invalid URL: {e}")))?;
let Some(host) = url.host_str() else {
return Ok(None);
};
let port = url.port_or_known_default().unwrap_or(443);
let addrs = resolve_and_validate(host, port)
.await
.map_err(|e| A2aError::Security(e.to_string()))?;
Ok(Some(PinnedTarget {
host: host.to_owned(),
addrs,
}))
}
fn needs_hardened_client(&self) -> bool {
self.security.require_tls || self.security.ssrf_protection
}
fn request_client(&self, pinned: Option<&PinnedTarget>) -> Result<reqwest::Client, A2aError> {
if self.needs_hardened_client() {
self.build_hardened_client(pinned)
} else {
Ok(self.client.clone())
}
}
fn build_hardened_client(
&self,
pinned: Option<&PinnedTarget>,
) -> Result<reqwest::Client, A2aError> {
let mut builder = reqwest::Client::builder()
.user_agent(concat!("zeph-a2a/", env!("CARGO_PKG_VERSION")))
.redirect(reqwest::redirect::Policy::none());
if self.security.require_tls {
builder = builder.https_only(true);
}
if let Some(target) = pinned {
builder = builder.resolve_to_addrs(&target.host, &target.addrs);
}
builder
.build()
.map_err(|e| A2aError::Security(format!("failed to build hardened client: {e}")))
}
#[tracing::instrument(name = "a2a.client.rpc_call", skip_all, err)]
async fn rpc_call<P: Serialize, R: DeserializeOwned>(
&self,
endpoint: &str,
method: &str,
params: P,
token: Option<&str>,
task_id: &str,
) -> Result<R, A2aError> {
let pinned = self.validate_endpoint(endpoint).await?;
let request_client = self.request_client(pinned.as_ref())?;
let request = JsonRpcRequest::new(method, params);
let mut req = request_client.post(endpoint).json(&request);
if let Some(t) = token {
req = req.bearer_auth(t);
}
if let Some(ibct_header) = self.ibct_header_value(endpoint, task_id) {
req = req.header("X-Zeph-IBCT", ibct_header);
}
let rpc_response: JsonRpcResponse<R> = tokio::time::timeout(self.request_timeout, async {
let resp = req.send().await?;
resp.json().await
})
.await
.map_err(|_| A2aError::Timeout(self.request_timeout))?
.map_err(A2aError::Http)?;
rpc_response.into_result().map_err(A2aError::from)
}
}
#[cfg(test)]
mod tests {
use std::assert_matches;
use std::net::IpAddr;
use super::*;
use zeph_common::net::is_private_ip;
use crate::jsonrpc::{JsonRpcError, JsonRpcResponse};
use crate::types::{
Artifact, Message, Part, Task, TaskArtifactUpdateEvent, TaskState, TaskStatus,
TaskStatusUpdateEvent,
};
#[test]
fn task_event_deserialize_status_update() {
let event = TaskStatusUpdateEvent {
kind: "status-update".into(),
task_id: "t-1".into(),
context_id: None,
status: TaskStatus {
state: TaskState::Working,
timestamp: "ts".into(),
message: Some(Message::user_text("thinking...")),
},
is_final: false,
};
let json = serde_json::to_string(&event).unwrap();
let parsed: TaskEvent = serde_json::from_str(&json).unwrap();
assert_matches!(parsed, TaskEvent::StatusUpdate(_));
}
#[test]
fn task_event_deserialize_artifact_update() {
let event = TaskArtifactUpdateEvent {
kind: "artifact-update".into(),
task_id: "t-1".into(),
context_id: None,
artifact: Artifact {
artifact_id: "a-1".into(),
name: None,
parts: vec![Part::text("result")],
metadata: None,
},
is_final: true,
};
let json = serde_json::to_string(&event).unwrap();
let parsed: TaskEvent = serde_json::from_str(&json).unwrap();
assert_matches!(parsed, TaskEvent::ArtifactUpdate(_));
}
#[test]
fn rpc_response_with_task_result() {
let task = Task {
id: "t-1".into(),
context_id: None,
status: TaskStatus {
state: TaskState::Completed,
timestamp: "ts".into(),
message: None,
},
artifacts: vec![],
history: vec![],
metadata: None,
};
let resp = JsonRpcResponse {
jsonrpc: "2.0".into(),
id: serde_json::Value::String("req-1".into()),
result: Some(task),
error: None,
};
let json = serde_json::to_string(&resp).unwrap();
let back: JsonRpcResponse<Task> = serde_json::from_str(&json).unwrap();
let task = back.into_result().unwrap();
assert_eq!(task.id, "t-1");
assert_eq!(task.status.state, TaskState::Completed);
}
#[test]
fn rpc_response_with_error() {
let resp: JsonRpcResponse<Task> = JsonRpcResponse {
jsonrpc: "2.0".into(),
id: serde_json::Value::String("req-1".into()),
result: None,
error: Some(JsonRpcError {
code: -32001,
message: "task not found".into(),
data: None,
}),
};
let json = serde_json::to_string(&resp).unwrap();
let back: JsonRpcResponse<Task> = serde_json::from_str(&json).unwrap();
let err = back.into_result().unwrap_err();
assert_eq!(err.code, -32001);
}
#[test]
fn a2a_client_construction() {
let client = A2aClient::new(reqwest::Client::new());
drop(client);
}
#[test]
fn is_private_ip_loopback() {
assert!(is_private_ip(IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)));
assert!(is_private_ip(IpAddr::V6(std::net::Ipv6Addr::LOCALHOST)));
}
#[test]
fn is_private_ip_private_ranges() {
assert!(is_private_ip("10.0.0.1".parse().unwrap()));
assert!(is_private_ip("172.16.0.1".parse().unwrap()));
assert!(is_private_ip("192.168.1.1".parse().unwrap()));
}
#[test]
fn is_private_ip_link_local() {
assert!(is_private_ip("169.254.0.1".parse().unwrap()));
}
#[test]
fn is_private_ip_unspecified() {
assert!(is_private_ip("0.0.0.0".parse().unwrap()));
assert!(is_private_ip("::".parse().unwrap()));
}
#[test]
fn is_private_ip_public() {
assert!(!is_private_ip("8.8.8.8".parse().unwrap()));
assert!(!is_private_ip("1.1.1.1".parse().unwrap()));
}
#[tokio::test]
async fn tls_enforcement_rejects_http() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: false,
});
let result = client.validate_endpoint("http://example.com/rpc").await;
assert!(result.is_err());
let err = result.unwrap_err();
assert_matches!(err, A2aError::Security(_));
assert!(err.to_string().contains("TLS required"));
}
#[tokio::test]
async fn tls_enforcement_allows_https() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: false,
});
let result = client.validate_endpoint("https://example.com/rpc").await;
assert!(result.is_ok());
}
#[tokio::test]
async fn ssrf_protection_rejects_localhost() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: false,
ssrf_protection: true,
});
let result = client.validate_endpoint("http://127.0.0.1:8080/rpc").await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("SSRF"));
}
#[tokio::test]
async fn no_security_allows_http_localhost() {
let client = A2aClient::new(reqwest::Client::new());
let result = client.validate_endpoint("http://127.0.0.1:8080/rpc").await;
assert!(result.is_ok());
}
#[test]
fn jsonrpc_request_serialization_for_send_message() {
let params = SendMessageParams {
message: Message::user_text("hello"),
configuration: None,
};
let req = JsonRpcRequest::new(METHOD_SEND_MESSAGE, params);
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("\"method\":\"message/send\""));
assert!(json.contains("\"jsonrpc\":\"2.0\""));
assert!(json.contains("\"hello\""));
}
#[test]
fn jsonrpc_request_serialization_for_get_task() {
let params = TaskIdParams {
id: "task-123".into(),
history_length: Some(5),
};
let req = JsonRpcRequest::new(METHOD_GET_TASK, params);
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("\"method\":\"tasks/get\""));
assert!(json.contains("\"task-123\""));
assert!(json.contains("\"historyLength\":5"));
}
#[test]
fn jsonrpc_request_serialization_for_cancel_task() {
let params = TaskIdParams {
id: "task-456".into(),
history_length: None,
};
let req = JsonRpcRequest::new(METHOD_CANCEL_TASK, params);
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("\"method\":\"tasks/cancel\""));
assert!(!json.contains("historyLength"));
}
#[test]
fn jsonrpc_request_serialization_for_stream() {
let params = SendMessageParams {
message: Message::user_text("stream me"),
configuration: None,
};
let req = JsonRpcRequest::new(METHOD_SEND_STREAMING_MESSAGE, params);
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("\"method\":\"message/stream\""));
}
#[tokio::test]
async fn send_message_connection_error() {
let client = A2aClient::new(reqwest::Client::new());
let params = SendMessageParams {
message: Message::user_text("hello"),
configuration: None,
};
let result = client
.send_message("http://127.0.0.1:1/rpc", params, None)
.await;
assert!(result.is_err());
assert_matches!(result.unwrap_err(), A2aError::Http(_));
}
#[tokio::test]
async fn get_task_connection_error() {
let client = A2aClient::new(reqwest::Client::new());
let params = TaskIdParams {
id: "t-1".into(),
history_length: None,
};
let result = client
.get_task("http://127.0.0.1:1/rpc", params, None)
.await;
assert!(result.is_err());
assert_matches!(result.unwrap_err(), A2aError::Http(_));
}
#[tokio::test]
async fn cancel_task_connection_error() {
let client = A2aClient::new(reqwest::Client::new());
let params = TaskIdParams {
id: "t-1".into(),
history_length: None,
};
let result = client
.cancel_task("http://127.0.0.1:1/rpc", params, None)
.await;
assert!(result.is_err());
assert_matches!(result.unwrap_err(), A2aError::Http(_));
}
#[tokio::test]
async fn stream_message_connection_error() {
let client = A2aClient::new(reqwest::Client::new());
let params = SendMessageParams {
message: Message::user_text("stream me"),
configuration: None,
};
let result = client
.stream_message("http://127.0.0.1:1/rpc", params, None)
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn stream_message_tls_required_rejects_http() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: false,
});
let params = SendMessageParams {
message: Message::user_text("hello"),
configuration: None,
};
let result = client
.stream_message("http://example.com/rpc", params, None)
.await;
match result {
Err(A2aError::Security(msg)) => assert!(msg.contains("TLS required")),
_ => panic!("expected Security error"),
}
}
#[tokio::test]
async fn send_message_tls_required_rejects_http() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: false,
});
let params = SendMessageParams {
message: Message::user_text("hello"),
configuration: None,
};
let result = client
.send_message("http://example.com/rpc", params, None)
.await;
assert!(result.is_err());
assert_matches!(result.unwrap_err(), A2aError::Security(_));
}
#[tokio::test]
async fn get_task_tls_required_rejects_http() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: false,
});
let params = TaskIdParams {
id: "t-1".into(),
history_length: None,
};
let result = client
.get_task("http://example.com/rpc", params, None)
.await;
assert!(result.is_err());
assert_matches!(result.unwrap_err(), A2aError::Security(_));
}
#[tokio::test]
async fn cancel_task_tls_required_rejects_http() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: false,
});
let params = TaskIdParams {
id: "t-1".into(),
history_length: None,
};
let result = client
.cancel_task("http://example.com/rpc", params, None)
.await;
assert!(result.is_err());
assert_matches!(result.unwrap_err(), A2aError::Security(_));
}
#[tokio::test]
async fn validate_endpoint_invalid_url_with_ssrf() {
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: false,
ssrf_protection: true,
});
let result = client.validate_endpoint("not-a-url").await;
assert!(result.is_err());
assert_matches!(result.unwrap_err(), A2aError::Security(_));
}
#[test]
fn with_security_returns_configured_client() {
let client =
A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy::hardened());
assert!(client.security.require_tls);
assert!(client.security.ssrf_protection);
}
#[test]
fn default_client_no_security() {
let client = A2aClient::new(reqwest::Client::new());
assert!(!client.security.require_tls);
assert!(!client.security.ssrf_protection);
}
#[test]
fn needs_hardened_client_reflects_policy() {
assert!(!A2aClient::new(reqwest::Client::new()).needs_hardened_client());
assert!(
A2aClient::new(reqwest::Client::new())
.with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: false,
})
.needs_hardened_client()
);
assert!(
A2aClient::new(reqwest::Client::new())
.with_security(SecurityPolicy {
require_tls: false,
ssrf_protection: true,
})
.needs_hardened_client()
);
assert!(
A2aClient::new(reqwest::Client::new())
.with_security(SecurityPolicy::hardened())
.needs_hardened_client()
);
}
#[test]
fn task_event_clone() {
let event = TaskEvent::StatusUpdate(TaskStatusUpdateEvent {
kind: "status-update".into(),
task_id: "t-1".into(),
context_id: None,
status: TaskStatus {
state: TaskState::Working,
timestamp: "ts".into(),
message: None,
},
is_final: false,
});
let cloned = event.clone();
let json1 = serde_json::to_string(&event).unwrap();
let json2 = serde_json::to_string(&cloned).unwrap();
assert_eq!(json1, json2);
}
#[test]
fn task_event_debug() {
let event = TaskEvent::ArtifactUpdate(TaskArtifactUpdateEvent {
kind: "artifact-update".into(),
task_id: "t-1".into(),
context_id: None,
artifact: Artifact {
artifact_id: "a-1".into(),
name: None,
parts: vec![Part::text("data")],
metadata: None,
},
is_final: true,
});
let dbg = format!("{event:?}");
assert!(dbg.contains("ArtifactUpdate"));
}
#[test]
fn is_private_ip_ipv4_non_private() {
assert!(!is_private_ip("93.184.216.34".parse().unwrap()));
}
#[test]
fn is_private_ip_ipv6_non_private() {
assert!(!is_private_ip("2001:db8::1".parse().unwrap()));
}
#[test]
fn rpc_response_error_takes_priority_over_result() {
let resp = JsonRpcResponse {
jsonrpc: "2.0".into(),
id: serde_json::Value::String("1".into()),
result: Some(Task {
id: "t-1".into(),
context_id: None,
status: TaskStatus {
state: TaskState::Completed,
timestamp: "ts".into(),
message: None,
},
artifacts: vec![],
history: vec![],
metadata: None,
}),
error: Some(JsonRpcError {
code: -32001,
message: "error".into(),
data: None,
}),
};
let err = resp.into_result().unwrap_err();
assert_eq!(err.code, -32001);
}
#[test]
fn rpc_response_neither_result_nor_error() {
let resp: JsonRpcResponse<Task> = JsonRpcResponse {
jsonrpc: "2.0".into(),
id: serde_json::Value::String("1".into()),
result: None,
error: None,
};
let err = resp.into_result().unwrap_err();
assert_eq!(err.code, -32603);
}
#[test]
fn with_ibct_key_sets_key_and_default_ttl() {
let key = IbctKey {
key_id: "k1".into(),
key_bytes: b"secret".to_vec(),
};
let client = A2aClient::new(reqwest::Client::new()).with_ibct_key(key);
assert!(client.ibct_key.is_some());
assert_eq!(client.ibct_ttl, Duration::from_mins(5));
}
#[test]
fn with_ibct_ttl_overrides_default() {
let key = IbctKey {
key_id: "k1".into(),
key_bytes: b"secret".to_vec(),
};
let client = A2aClient::new(reqwest::Client::new())
.with_ibct_key(key)
.with_ibct_ttl(Duration::from_mins(1));
assert_eq!(client.ibct_ttl, Duration::from_mins(1));
}
#[test]
fn no_ibct_key_configured_yields_no_header() {
let client = A2aClient::new(reqwest::Client::new());
assert!(
client
.ibct_header_value("https://agent.example.com", "task-1")
.is_none()
);
}
#[cfg(feature = "ibct")]
#[test]
fn ibct_key_configured_yields_encoded_header() {
let key = IbctKey {
key_id: "k1".into(),
key_bytes: b"secret".to_vec(),
};
let client = A2aClient::new(reqwest::Client::new()).with_ibct_key(key);
let header = client
.ibct_header_value("https://agent.example.com", "task-1")
.expect("header should be issued when ibct feature is enabled");
let decoded = crate::ibct::Ibct::decode(&header).unwrap();
assert_eq!(decoded.task_id, "task-1");
assert_eq!(decoded.endpoint, "https://agent.example.com");
}
#[test]
fn ibct_scope_origin_strips_path_and_query() {
assert_eq!(
ibct_scope_origin("https://agent.example.com/a2a"),
"https://agent.example.com"
);
assert_eq!(
ibct_scope_origin("https://agent.example.com/a2a/stream"),
"https://agent.example.com"
);
assert_eq!(
ibct_scope_origin("http://127.0.0.1:8080/a2a/stream?x=1"),
"http://127.0.0.1:8080"
);
}
#[test]
fn ibct_scope_origin_agrees_across_both_a2a_routes() {
let a2a = ibct_scope_origin("http://127.0.0.1:8080/a2a");
let stream = ibct_scope_origin("http://127.0.0.1:8080/a2a/stream");
assert_eq!(
a2a, stream,
"both routes on the same agent must scope to the same IBCT endpoint"
);
}
#[test]
fn ibct_scope_origin_falls_back_to_input_on_unparseable_url() {
assert_eq!(ibct_scope_origin("not-a-url"), "not-a-url");
}
#[cfg(feature = "ibct")]
#[test]
fn ibct_header_value_scopes_token_to_origin_not_full_path() {
let key = IbctKey {
key_id: "k1".into(),
key_bytes: b"secret".to_vec(),
};
let client = A2aClient::new(reqwest::Client::new()).with_ibct_key(key);
let header_a2a = client
.ibct_header_value("http://127.0.0.1:8080/a2a", "task-1")
.unwrap();
let header_stream = client
.ibct_header_value("http://127.0.0.1:8080/a2a/stream", "task-1")
.unwrap();
let decoded_a2a = crate::ibct::Ibct::decode(&header_a2a).unwrap();
let decoded_stream = crate::ibct::Ibct::decode(&header_stream).unwrap();
assert_eq!(decoded_a2a.endpoint, "http://127.0.0.1:8080");
assert_eq!(
decoded_a2a.endpoint, decoded_stream.endpoint,
"a token issued for /a2a and one issued for /a2a/stream on the same agent must \
carry the same endpoint scope, since the server verifies both against one \
pathless card.url"
);
}
#[test]
fn task_event_serialize_round_trip() {
let event = TaskEvent::StatusUpdate(TaskStatusUpdateEvent {
kind: "status-update".into(),
task_id: "t-1".into(),
context_id: Some("ctx-1".into()),
status: TaskStatus {
state: TaskState::Completed,
timestamp: "2025-01-01T00:00:00Z".into(),
message: Some(Message::user_text("done")),
},
is_final: true,
});
let json = serde_json::to_string(&event).unwrap();
let back: TaskEvent = serde_json::from_str(&json).unwrap();
assert_matches!(back, TaskEvent::StatusUpdate(_));
}
}
#[cfg(test)]
mod wiremock_tests {
use std::assert_matches;
use tokio_stream::StreamExt;
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use crate::client::{A2aClient, PinnedTarget, SecurityPolicy};
use crate::jsonrpc::{SendMessageParams, TaskIdParams};
use crate::testing::*;
use crate::types::Message;
#[tokio::test]
async fn send_message_success() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.respond_with(task_rpc_response("task-1", "submitted"))
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new());
let params = SendMessageParams {
message: Message::user_text("hello"),
configuration: None,
};
let task = client
.send_message(&format!("{}/rpc", server.uri()), params, None)
.await
.unwrap();
assert_eq!(task.id, "task-1");
}
#[tokio::test]
async fn send_message_rpc_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.respond_with(task_rpc_error_response(-32001, "task not found"))
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new());
let params = SendMessageParams {
message: Message::user_text("hi"),
configuration: None,
};
let result = client
.send_message(&format!("{}/rpc", server.uri()), params, None)
.await;
assert!(result.is_err());
let err = result.unwrap_err();
assert_matches!(err, crate::error::A2aError::JsonRpc { code: -32001, .. });
}
#[tokio::test]
async fn send_message_with_bearer_auth() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.and(header("authorization", "Bearer secret-token"))
.respond_with(task_rpc_response("task-auth", "submitted"))
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new());
let params = SendMessageParams {
message: Message::user_text("secure"),
configuration: None,
};
let task = client
.send_message(
&format!("{}/rpc", server.uri()),
params,
Some("secret-token"),
)
.await
.unwrap();
assert_eq!(task.id, "task-auth");
}
#[tokio::test]
async fn get_task_success() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.respond_with(task_rpc_response("task-get", "completed"))
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new());
let params = TaskIdParams {
id: "task-get".into(),
history_length: None,
};
let task = client
.get_task(&format!("{}/rpc", server.uri()), params, None)
.await
.unwrap();
assert_eq!(task.id, "task-get");
}
#[tokio::test]
async fn cancel_task_success() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.respond_with(task_rpc_response("task-cancel", "canceled"))
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new());
let params = TaskIdParams {
id: "task-cancel".into(),
history_length: None,
};
let task = client
.cancel_task(&format!("{}/rpc", server.uri()), params, None)
.await
.unwrap();
assert_eq!(task.id, "task-cancel");
}
#[tokio::test]
async fn stream_message_success() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.respond_with(sse_task_events_response("task-stream", "result content"))
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new());
let params = SendMessageParams {
message: Message::user_text("stream"),
configuration: None,
};
let stream = client
.stream_message(&format!("{}/rpc", server.uri()), params, None)
.await
.unwrap();
let events: Vec<_> = stream.collect().await;
assert!(!events.is_empty());
}
#[tokio::test]
async fn stream_message_http_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.respond_with(ResponseTemplate::new(500).set_body_string("Internal Server Error"))
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new());
let params = SendMessageParams {
message: Message::user_text("fail"),
configuration: None,
};
let result = client
.stream_message(&format!("{}/rpc", server.uri()), params, None)
.await;
let err = result.err().expect("expected error");
assert_matches!(err, crate::error::A2aError::Stream(_));
}
#[tokio::test]
async fn rpc_call_times_out() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(std::time::Duration::from_secs(5))
.set_body_json(serde_json::json!({
"jsonrpc": "2.0",
"id": "req-1",
"result": {
"id": "t-1",
"status": {"state": "completed", "timestamp": "2026-01-01T00:00:00Z"}
}
})),
)
.mount(&server)
.await;
let client = A2aClient::new(reqwest::Client::new())
.with_request_timeout(std::time::Duration::from_millis(100));
let params = SendMessageParams {
message: Message::user_text("hello"),
configuration: None,
};
let result = client
.send_message(&format!("{}/rpc", server.uri()), params, None)
.await;
assert!(result.is_err());
assert!(
matches!(result.unwrap_err(), crate::error::A2aError::Timeout(_)),
"expected Timeout error"
);
}
#[tokio::test]
async fn hardened_client_pins_connection_bypassing_dns() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.respond_with(ResponseTemplate::new(200).set_body_string("ok"))
.mount(&server)
.await;
let addr = *server.address();
let fake_host = "zeph-a2a-pin-test.invalid";
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: false,
ssrf_protection: true,
});
let pinned = PinnedTarget {
host: fake_host.to_owned(),
addrs: vec![addr],
};
let hardened = client.build_hardened_client(Some(&pinned)).unwrap();
let resp = hardened
.get(format!("http://{fake_host}:{}/", addr.port()))
.send()
.await
.unwrap_or_else(|e| panic!("pinned request to unresolvable host failed: {e}"));
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "ok");
}
#[tokio::test]
async fn hardened_client_does_not_auto_follow_redirect_to_private_ip() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.respond_with(
ResponseTemplate::new(302).insert_header("Location", "http://127.0.0.1:9/internal"),
)
.mount(&server)
.await;
let addr = *server.address();
let fake_host = "zeph-a2a-redirect-test.invalid";
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: false,
ssrf_protection: true,
});
let pinned = PinnedTarget {
host: fake_host.to_owned(),
addrs: vec![addr],
};
let hardened = client.build_hardened_client(Some(&pinned)).unwrap();
let resp = hardened
.get(format!("http://{fake_host}:{}/start", addr.port()))
.send()
.await
.unwrap();
assert_eq!(resp.status(), reqwest::StatusCode::FOUND);
assert_eq!(
resp.headers().get(reqwest::header::LOCATION).unwrap(),
"http://127.0.0.1:9/internal"
);
}
#[tokio::test]
async fn hardened_client_with_require_tls_rejects_plaintext_connection() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.respond_with(ResponseTemplate::new(200))
.mount(&server)
.await;
let addr = *server.address();
let fake_host = "zeph-a2a-tls-test.invalid";
let client = A2aClient::new(reqwest::Client::new()).with_security(SecurityPolicy {
require_tls: true,
ssrf_protection: true,
});
let pinned = PinnedTarget {
host: fake_host.to_owned(),
addrs: vec![addr],
};
let hardened = client.build_hardened_client(Some(&pinned)).unwrap();
let result = hardened
.get(format!("http://{fake_host}:{}/", addr.port()))
.send()
.await;
assert!(
result.is_err(),
"https_only(true) must reject a plain http:// URL"
);
}
#[cfg(feature = "ibct")]
#[tokio::test]
async fn send_message_attaches_ibct_header_when_configured() {
use wiremock::matchers::header_exists;
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/rpc"))
.and(header_exists("x-zeph-ibct"))
.respond_with(task_rpc_response("task-ibct", "submitted"))
.mount(&server)
.await;
let key = crate::ibct::IbctKey {
key_id: "k1".into(),
key_bytes: b"secret".to_vec(),
};
let client = A2aClient::new(reqwest::Client::new()).with_ibct_key(key);
let params = SendMessageParams {
message: Message::user_text("hello"),
configuration: None,
};
let task = client
.send_message(&format!("{}/rpc", server.uri()), params, None)
.await
.unwrap();
assert_eq!(task.id, "task-ibct");
}
}