use wiremock::matchers::{body_string_contains, header, method};
use wiremock::{Mock, MockServer, Request, ResponseTemplate};
use yoagent::mcp::transport::McpTransport;
use yoagent::mcp::types::JsonRpcRequest;
use yoagent::mcp::{HttpTransport, McpClient};
struct HeaderAbsent(&'static str);
impl wiremock::Match for HeaderAbsent {
fn matches(&self, request: &Request) -> bool {
!request.headers.contains_key(self.0)
}
}
struct HeaderContains(&'static str, &'static str);
impl wiremock::Match for HeaderContains {
fn matches(&self, request: &Request) -> bool {
request
.headers
.get(self.0)
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.contains(self.1))
}
}
fn sse(payload: &str) -> String {
format!("event: message\ndata: {payload}\n\n")
}
async fn mount_body(server: &MockServer, body: String, content_type: &str) {
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_raw(body, content_type))
.mount(server)
.await;
}
#[tokio::test]
async fn plain_json_response_still_parses() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("ping", None);
mount_body(
&server,
format!(
r#"{{"jsonrpc":"2.0","id":{},"result":{{"ok":true}}}}"#,
request.id
),
"application/json",
)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport
.send(request)
.await
.expect("plain JSON must parse");
assert_eq!(response.result.unwrap()["ok"], true);
}
#[tokio::test]
async fn sse_framed_response_parses() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/list", None);
mount_body(
&server,
sse(&format!(
r#"{{"jsonrpc":"2.0","id":{},"result":{{"tools":[]}}}}"#,
request.id
)),
"text/event-stream",
)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport.send(request).await.expect("SSE must parse");
assert!(response.result.unwrap()["tools"].is_array());
}
#[tokio::test]
async fn bom_prefixed_plain_json_body_still_parses() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("ping", None);
mount_body(
&server,
format!(
"\u{feff}{{\"jsonrpc\":\"2.0\",\"id\":{},\"result\":{{\"ok\":true}}}}",
request.id
),
"application/json",
)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport
.send(request)
.await
.expect("a BOM must not hide the response");
assert_eq!(response.result.unwrap()["ok"], true);
}
#[tokio::test]
async fn non_utf8_charset_is_refused_rather_than_mangled() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("ping", None);
mount_body(
&server,
format!(r#"{{"jsonrpc":"2.0","id":{},"result":{{}}}}"#, request.id),
"application/json; charset=iso-8859-1",
)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let err = transport
.send(request)
.await
.expect_err("a non-UTF-8 charset must not be silently decoded as UTF-8");
assert!(err.to_string().contains("charset"), "got: {err}");
}
#[tokio::test]
async fn frame_with_both_method_and_result_is_not_the_response() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/list", None);
let body = format!(
concat!(
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{id},\"method\":\"notifications/progress\",\"result\":{{\"bogus\":true}}}}\n\n",
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"real\":true}}}}\n\n"
),
id = request.id
);
mount_body(&server, body, "text/event-stream").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport.send(request).await.expect("send");
assert_eq!(
response.result.expect("the well-formed response must win")["real"],
true
);
}
#[tokio::test]
async fn sse_skips_frames_that_are_not_json_rpc() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/list", None);
let body = format!(
concat!(
": keep-alive comment\n\n",
"event: ping\ndata: {{\"unrelated\":true}}\n\n",
"id: 42\nevent: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{},\"result\":{{\"found\":true}}}}\n\n"
),
request.id
);
mount_body(&server, body, "text/event-stream").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport
.send(request)
.await
.expect("payload must be found");
assert_eq!(response.result.unwrap()["found"], true);
}
#[tokio::test]
async fn notification_frames_do_not_shadow_the_response() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/list", None);
let body = format!(
concat!(
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"method\":\"notifications/message\",\"params\":{{\"level\":\"info\",\"data\":\"searching\"}}}}\n\n",
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\",\"params\":{{\"progress\":1}}}}\n\n",
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{},\"result\":{{\"tools\":[{{\"name\":\"web_search\"}}]}}}}\n\n"
),
request.id
);
mount_body(&server, body, "text/event-stream").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport.send(request).await.expect("send");
let result = response
.result
.expect("the result must win over preceding notification frames");
assert_eq!(result["tools"][0]["name"], "web_search");
}
#[tokio::test]
async fn server_initiated_request_does_not_shadow_the_response() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/call", None);
let body = format!(
concat!(
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":77,\"method\":\"sampling/createMessage\",\"params\":{{}}}}\n\n",
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{},\"result\":{{\"content\":[]}}}}\n\n"
),
request.id
);
mount_body(&server, body, "text/event-stream").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport.send(request).await.expect("send");
assert!(
response.result.is_some(),
"the real result must be selected"
);
}
#[tokio::test]
async fn bare_ack_frame_does_not_shadow_the_response() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/list", None);
let body = format!(
concat!(
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{id}}}\n\n",
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{id},\"result\":{{\"real\":true}}}}\n\n"
),
id = request.id
);
mount_body(&server, body, "text/event-stream").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport.send(request).await.expect("send");
assert_eq!(
response
.result
.expect("the real result must win over a bare ack")["real"],
true
);
}
#[tokio::test]
async fn json_rpc_error_over_sse_reaches_the_caller() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/call", None);
let body = format!(
concat!(
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"method\":\"notifications/message\",\"params\":{{}}}}\n\n",
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\"id\":{},\"error\":{{\"code\":-32602,\"message\":\"path must be absolute\"}}}}\n\n"
),
request.id
);
mount_body(&server, body, "text/event-stream").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport.send(request).await.expect("send");
let error = response.error.expect("the server's error must survive");
assert_eq!(error.code, -32602);
assert!(error.message.contains("absolute"));
}
#[tokio::test]
async fn mismatched_response_id_is_rejected() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/list", None);
mount_body(
&server,
sse(r#"{"jsonrpc":"2.0","id":999999,"result":{"stale":true}}"#),
"text/event-stream",
)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let err = transport
.send(request)
.await
.expect_err("a response for another id must not be returned as ours");
assert!(
err.to_string().contains("no JSON-RPC response"),
"got: {err}"
);
}
#[tokio::test]
async fn multi_line_data_frames_are_joined() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("tools/list", None);
let body = format!(
"event: message\ndata: {{\"jsonrpc\":\"2.0\",\ndata: \"id\":{},\ndata: \"result\":{{\"joined\":true}}}}\n\n",
request.id
);
mount_body(&server, body, "text/event-stream").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let response = transport.send(request).await.expect("send");
assert_eq!(response.result.unwrap()["joined"], true);
}
#[tokio::test]
async fn session_id_is_captured_then_replayed() {
let server = MockServer::start().await;
let first = JsonRpcRequest::new("initialize", None);
let second = JsonRpcRequest::new("tools/list", None);
Mock::given(method("POST"))
.and(HeaderAbsent("mcp-session-id"))
.and(body_string_contains("\"method\":\"initialize\""))
.respond_with(
ResponseTemplate::new(200)
.insert_header("Mcp-Session-Id", "sess-abc123")
.set_body_raw(
sse(&format!(
r#"{{"jsonrpc":"2.0","id":{},"result":{{"initialized":true}}}}"#,
first.id
)),
"text/event-stream",
),
)
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(header("mcp-session-id", "sess-abc123"))
.and(body_string_contains("tools/list"))
.respond_with(ResponseTemplate::new(200).set_body_raw(
sse(&format!(
r#"{{"jsonrpc":"2.0","id":{},"result":{{"tools":[]}}}}"#,
second.id
)),
"text/event-stream",
))
.expect(1)
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
transport.send(first).await.expect("initialize");
transport
.send(second)
.await
.expect("follow-up must reuse the session");
}
#[tokio::test]
async fn expired_session_is_cleared_and_named() {
let server = MockServer::start().await;
let first = JsonRpcRequest::new("initialize", None);
Mock::given(method("POST"))
.and(HeaderAbsent("mcp-session-id"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("Mcp-Session-Id", "sess-dead")
.set_body_raw(
format!(r#"{{"jsonrpc":"2.0","id":{},"result":{{}}}}"#, first.id),
"application/json",
),
)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(header("mcp-session-id", "sess-dead"))
.respond_with(ResponseTemplate::new(404))
.expect(1)
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
transport.send(first).await.expect("initialize");
let err = transport
.send(JsonRpcRequest::new("tools/list", None))
.await
.expect_err("404 must surface");
assert!(
err.to_string().contains("session expired"),
"the error must name the cause, got: {err}"
);
let _ = transport
.send(JsonRpcRequest::new("initialize", None))
.await;
let last = server
.received_requests()
.await
.unwrap()
.pop()
.expect("a third request was sent");
assert!(
!last.headers.contains_key("mcp-session-id"),
"the expired session must not be replayed"
);
}
#[tokio::test]
async fn accepted_with_empty_body_is_not_an_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(202))
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let request = JsonRpcRequest::new("notifications/initialized", None);
let request_id = request.id;
let response = transport
.send(request)
.await
.expect("202 with no body must not be an error");
assert_eq!(response.id, Some(request_id));
assert!(response.result.is_none() && response.error.is_none());
}
#[tokio::test]
async fn empty_body_on_plain_200_is_an_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200))
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let err = transport
.send(JsonRpcRequest::new("tools/list", None))
.await
.expect_err("an empty 200 must not be dressed up as success");
assert!(err.to_string().contains("empty body"), "got: {err}");
}
#[tokio::test]
async fn accept_header_advertises_both_framings() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("ping", None);
Mock::given(method("POST"))
.and(HeaderContains("accept", "application/json"))
.and(HeaderContains("accept", "text/event-stream"))
.respond_with(ResponseTemplate::new(200).set_body_raw(
format!(r#"{{"jsonrpc":"2.0","id":{},"result":{{}}}}"#, request.id),
"application/json",
))
.expect(1)
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
transport.send(request).await.expect("send");
}
#[tokio::test]
async fn non_2xx_preserves_the_server_explanation() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(400).set_body_raw(
r#"{"error":{"message":"Mcp-Session-Id header required"}}"#,
"application/json",
))
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let err = transport
.send(JsonRpcRequest::new("tools/list", None))
.await
.expect_err("400 must surface");
assert!(
err.to_string().contains("Mcp-Session-Id header required"),
"the server's explanation must survive, got: {err}"
);
}
#[tokio::test]
async fn close_deletes_the_session_when_one_exists() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("initialize", None);
Mock::given(method("POST"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("Mcp-Session-Id", "sess-xyz")
.set_body_raw(
format!(r#"{{"jsonrpc":"2.0","id":{},"result":{{}}}}"#, request.id),
"application/json",
),
)
.mount(&server)
.await;
Mock::given(method("DELETE"))
.and(header("mcp-session-id", "sess-xyz"))
.respond_with(ResponseTemplate::new(204))
.expect(1)
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
transport.send(request).await.expect("initialize");
transport.close().await.expect("close");
}
#[tokio::test]
async fn close_without_a_session_sends_nothing() {
let server = MockServer::start().await;
let transport = HttpTransport::new(&server.uri()).unwrap();
transport.close().await.expect("close must be a no-op");
assert!(
server.received_requests().await.unwrap().is_empty(),
"close() without a session must not touch the network"
);
}
#[tokio::test]
async fn close_ignores_a_server_that_rejects_delete() {
let server = MockServer::start().await;
let request = JsonRpcRequest::new("initialize", None);
Mock::given(method("POST"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("Mcp-Session-Id", "sess-xyz")
.set_body_raw(
format!(r#"{{"jsonrpc":"2.0","id":{},"result":{{}}}}"#, request.id),
"application/json",
),
)
.mount(&server)
.await;
Mock::given(method("DELETE"))
.respond_with(ResponseTemplate::new(405))
.mount(&server)
.await;
let transport = HttpTransport::new(&server.uri()).unwrap();
transport.send(request).await.expect("initialize");
transport
.close()
.await
.expect("a 405 on DELETE must not fail close()");
}
#[tokio::test]
async fn unparseable_body_is_an_error_with_context() {
let server = MockServer::start().await;
mount_body(&server, "<html>gateway</html>".into(), "text/html").await;
let transport = HttpTransport::new(&server.uri()).unwrap();
let err = transport
.send(JsonRpcRequest::new("ping", None))
.await
.expect_err("an unparseable body must be an error");
assert!(
err.to_string().contains("gateway"),
"error should quote the body, got: {err}"
);
}
#[tokio::test]
async fn e2e_handshake_over_streamable_http() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(body_string_contains("\"method\":\"initialize\""))
.respond_with(
ResponseTemplate::new(200)
.insert_header("Mcp-Session-Id", "sess-e2e")
.set_body_raw(
sse(r#"{"jsonrpc":"2.0","result":{"protocolVersion":"2024-11-05","capabilities":{},"serverInfo":{"name":"probe","version":"1.0"}}}"#),
"text/event-stream",
),
)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(body_string_contains("notifications/initialized"))
.respond_with(ResponseTemplate::new(202))
.mount(&server)
.await;
Mock::given(method("POST"))
.and(body_string_contains("tools/list"))
.and(header("mcp-session-id", "sess-e2e"))
.respond_with(ResponseTemplate::new(200).set_body_raw(
concat!(
"event: message\ndata: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/message\",\"params\":{\"level\":\"info\"}}\n\n",
"event: message\ndata: {\"jsonrpc\":\"2.0\",\"result\":{\"tools\":[{\"name\":\"web_search\",\"inputSchema\":{}}]}}\n\n"
),
"text/event-stream",
))
.mount(&server)
.await;
let client = McpClient::connect_http(&server.uri())
.await
.expect("handshake over SSE must succeed");
assert_eq!(client.server_info().unwrap().name, "probe");
let tools = client.list_tools().await.expect("tools/list over SSE");
assert_eq!(tools[0].name, "web_search");
}