1use std::convert::Infallible;
21use std::net::SocketAddr;
22use std::sync::Arc;
23
24use axum::Json;
25use axum::extract::State;
26use axum::http::{HeaderMap, HeaderName, StatusCode, header};
27use axum::response::sse::{Event, Sse};
28use axum::response::{IntoResponse, Response};
29use axum::routing::post;
30use serde_json::Value;
31
32use crate::mcp::McpServer;
33use crate::mcp::server::{
34 SUPPORTED_PROTOCOL_VERSIONS, request_protocol_version, shape_modern_result,
35};
36
37#[derive(Clone)]
39pub struct HttpTransport {
40 server: Arc<McpServer>,
41}
42
43impl HttpTransport {
44 pub fn new(server: Arc<McpServer>) -> Self {
46 Self { server }
47 }
48
49 pub fn router(&self) -> axum::Router {
51 axum::Router::new()
52 .route("/mcp", post(rpc_handler))
53 .with_state(self.clone())
54 }
55}
56
57pub async fn serve(server: Arc<McpServer>, addr: SocketAddr) -> std::io::Result<()> {
59 let transport = HttpTransport::new(server);
60 let listener = tokio::net::TcpListener::bind(addr).await?;
61 axum::serve(listener, transport.router()).await?;
62 Ok(())
63}
64
65const ERR_INVALID_REQUEST: i32 = -32600;
67const ERR_METHOD_NOT_FOUND: i32 = -32601;
69const ERR_INVALID_PARAMS: i32 = -32602;
71const ERR_UNAUTHORIZED: i32 = -32001;
73const ERR_HEADER_MISMATCH: i32 = -32020;
75const ERR_UNSUPPORTED_VERSION: i32 = -32022;
77
78const MODERN_METHODS: &[&str] = &[
81 "server/discover",
82 "tools/list",
83 "tools/call",
84 "resources/list",
85 "resources/templates/list",
86 "resources/read",
87 "prompts/list",
88 "prompts/get",
89 "subscriptions/listen",
90];
91
92async fn dispatch_async(server: &McpServer, req: &Value) -> Option<Value> {
100 let id = req.get("id")?.clone();
102 let method = req.get("method").and_then(|v| v.as_str()).unwrap_or("");
103
104 let modern = if method == "initialize" {
107 None
108 } else {
109 request_protocol_version(req)
110 };
111 if let Some(version) = modern {
112 if !SUPPORTED_PROTOCOL_VERSIONS.contains(&version) {
113 return Some(rpc_error(
114 Some(id),
115 ERR_UNSUPPORTED_VERSION,
116 "Unsupported protocol version",
117 Some(serde_json::json!({
118 "supported": SUPPORTED_PROTOCOL_VERSIONS,
119 "requested": version,
120 })),
121 ));
122 }
123 match method {
124 "server/discover" => {
125 let mut result = server.discover_response();
126 shape_modern_result(version, method, &mut result);
127 return Some(rpc_result(id, result));
128 }
129 "ping" => {
131 return Some(rpc_error(
132 Some(id),
133 ERR_METHOD_NOT_FOUND,
134 "Method not found: ping (removed in protocol 2026-07-28)",
135 None,
136 ));
137 }
138 "subscriptions/listen" => {
141 let (_, close) = server.subscription_ack_and_close(&id);
142 return Some(rpc_result(id, close));
143 }
144 _ => {}
145 }
146 }
147
148 let result: Result<Value, (i32, String)> = match method {
149 "initialize" => {
150 let requested = req
151 .pointer("/params/protocolVersion")
152 .and_then(|v| v.as_str());
153 Ok(server.initialize_response(requested))
154 }
155 "ping" => Ok(serde_json::json!({})),
156 "tools/list" => Ok(serde_json::json!({ "tools": server.tools() })),
157 "resources/list" => Ok(serde_json::json!({ "resources": server.resources() })),
158 "resources/templates/list" => Ok(serde_json::json!({ "resourceTemplates": [] })),
159 "prompts/list" => Ok(serde_json::json!({ "prompts": server.prompts() })),
160 "prompts/get" => {
161 let name = req
162 .pointer("/params/name")
163 .and_then(|v| v.as_str())
164 .unwrap_or("");
165 let args = req
166 .pointer("/params/arguments")
167 .cloned()
168 .unwrap_or(serde_json::json!({}));
169 server
170 .get_prompt(name, args)
171 .map_err(|e| (ERR_INVALID_PARAMS, e.to_string()))
172 }
173 "resources/read" => {
174 let uri = req
175 .pointer("/params/uri")
176 .and_then(|v| v.as_str())
177 .unwrap_or("");
178 server
179 .read_resource(uri, serde_json::json!({}))
180 .map(|content| {
181 serde_json::json!({
182 "contents": [{ "uri": uri, "text": content.to_string() }]
183 })
184 })
185 .map_err(|e| (ERR_INVALID_PARAMS, e.to_string()))
186 }
187 "tools/call" => {
188 let name = req
189 .pointer("/params/name")
190 .and_then(|v| v.as_str())
191 .unwrap_or("");
192 let params = req
193 .pointer("/params/arguments")
194 .cloned()
195 .unwrap_or(serde_json::json!(null));
196 if !server.has_tool(name) {
197 Err((ERR_INVALID_PARAMS, format!("Unknown tool: {name}")))
198 } else if let Err(e) = server.validate_tool_args(name, ¶ms) {
199 Err((ERR_INVALID_PARAMS, e))
200 } else {
201 match server.call_tool_async(name, params).await {
203 Ok(r) => Ok(serde_json::json!({
204 "content": [{ "type": "text", "text": r.to_string() }],
205 "isError": false
206 })),
207 Err(e) => Ok(serde_json::json!({
208 "content": [{ "type": "text", "text": e.to_string() }],
209 "isError": true
210 })),
211 }
212 }
213 }
214 _ => Err((ERR_METHOD_NOT_FOUND, format!("Method not found: {method}"))),
215 };
216
217 Some(match result {
218 Ok(mut value) => {
219 if let Some(version) = modern {
220 shape_modern_result(version, method, &mut value);
221 }
222 rpc_result(id, value)
223 }
224 Err((code, message)) => rpc_error(Some(id), code, &message, None),
225 })
226}
227
228fn rpc_result(id: Value, result: Value) -> Value {
230 serde_json::json!({ "jsonrpc": "2.0", "id": id, "result": result })
231}
232
233fn rpc_error(id: Option<Value>, code: i32, message: &str, data: Option<Value>) -> Value {
237 let mut error = serde_json::json!({ "code": code, "message": message });
238 if let Some(d) = data {
239 error["data"] = d;
240 }
241 serde_json::json!({ "jsonrpc": "2.0", "id": id, "error": error })
242}
243
244fn authorized(server: &McpServer, headers: &HeaderMap) -> bool {
247 let auth = headers
248 .get("authorization")
249 .and_then(|v| v.to_str().ok())
250 .unwrap_or("");
251 server.check_auth(auth)
252}
253
254fn origin_allowed(headers: &HeaderMap) -> bool {
260 let Some(origin) = headers.get("origin").and_then(|v| v.to_str().ok()) else {
261 return true; };
263 if origin == "null" {
264 return false;
265 }
266 origin
270 .split_once("://")
271 .map(|(_, host_port)| {
272 if let Some(rest) = host_port.strip_prefix('[') {
273 rest.split(']').next().unwrap_or("") } else {
275 host_port.split(':').next().unwrap_or("")
276 }
277 })
278 .is_some_and(|host| host == "localhost" || host == "127.0.0.1" || host == "::1")
279}
280
281fn forbidden_response(id: Option<Value>) -> Json<Value> {
282 Json(serde_json::json!({
283 "jsonrpc": "2.0",
284 "id": id,
285 "error": { "code": ERR_UNAUTHORIZED, "message": "Forbidden origin" }
286 }))
287}
288
289fn unauthorized_response(id: Option<Value>) -> Json<Value> {
290 Json(serde_json::json!({
291 "jsonrpc": "2.0",
292 "id": id,
293 "error": { "code": ERR_UNAUTHORIZED, "message": "Unauthorized" }
294 }))
295}
296
297async fn dispatch_any(server: &McpServer, req: &Value) -> Option<Value> {
300 let Some(batch) = req.as_array() else {
301 return dispatch_async(server, req).await;
302 };
303 let mut out = Vec::with_capacity(batch.len());
304 for item in batch {
305 if let Some(resp) = dispatch_async(server, item).await {
306 out.push(resp);
307 }
308 }
309 if out.is_empty() {
310 None
311 } else {
312 Some(Value::Array(out))
313 }
314}
315
316fn decode_header_value(raw: &str) -> String {
320 if let Some(encoded) = raw
321 .strip_prefix("=?base64?")
322 .and_then(|s| s.strip_suffix("?="))
323 {
324 use base64::Engine as _;
325 if let Ok(decoded) = base64::engine::general_purpose::STANDARD.decode(encoded)
326 && let Ok(text) = String::from_utf8(decoded)
327 {
328 return text;
329 }
330 }
331 raw.to_string()
332}
333
334fn validate_modern_headers(
341 req: &Value,
342 headers: &HeaderMap,
343) -> Result<(), (StatusCode, Json<Value>)> {
344 let id = req.get("id").cloned();
345 let reject = |status: StatusCode, code: i32, message: &str, data: Option<Value>| {
346 Err((status, Json(rpc_error(id, code, message, data))))
347 };
348 let version = request_protocol_version(req).expect("caller checked for _meta version");
349 let method = req.get("method").and_then(|v| v.as_str()).unwrap_or("");
350
351 if !SUPPORTED_PROTOCOL_VERSIONS.contains(&version) {
352 return reject(
353 StatusCode::BAD_REQUEST,
354 ERR_UNSUPPORTED_VERSION,
355 "Unsupported protocol version",
356 Some(serde_json::json!({
357 "supported": SUPPORTED_PROTOCOL_VERSIONS,
358 "requested": version,
359 })),
360 );
361 }
362 match headers
363 .get("mcp-protocol-version")
364 .and_then(|v| v.to_str().ok())
365 {
366 None => {
367 return reject(
368 StatusCode::BAD_REQUEST,
369 ERR_HEADER_MISMATCH,
370 "Missing required MCP-Protocol-Version header",
371 None,
372 );
373 }
374 Some(h) if h != version => {
375 return reject(
376 StatusCode::BAD_REQUEST,
377 ERR_HEADER_MISMATCH,
378 "MCP-Protocol-Version header does not match the body _meta version",
379 None,
380 );
381 }
382 _ => {}
383 }
384 match headers.get("mcp-method").and_then(|v| v.to_str().ok()) {
385 None => {
386 return reject(
387 StatusCode::BAD_REQUEST,
388 ERR_HEADER_MISMATCH,
389 "Missing required Mcp-Method header",
390 None,
391 );
392 }
393 Some(h) if h != method => {
394 return reject(
395 StatusCode::BAD_REQUEST,
396 ERR_HEADER_MISMATCH,
397 &format!("Mcp-Method header '{h}' does not match body method '{method}'"),
398 None,
399 );
400 }
401 _ => {}
402 }
403 let name_source = match method {
407 "tools/call" | "prompts/get" => req.pointer("/params/name"),
408 "resources/read" => req.pointer("/params/uri"),
409 _ => None,
410 };
411 if let Some(source) = name_source
412 && !source.is_null()
413 {
414 match headers.get("mcp-name").and_then(|v| v.to_str().ok()) {
415 None => {
416 return reject(
417 StatusCode::BAD_REQUEST,
418 ERR_HEADER_MISMATCH,
419 "Missing required Mcp-Name header",
420 None,
421 );
422 }
423 Some(h) if decode_header_value(h) != source.as_str().unwrap_or("") => {
424 return reject(
425 StatusCode::BAD_REQUEST,
426 ERR_HEADER_MISMATCH,
427 "Mcp-Name header does not match the body value",
428 None,
429 );
430 }
431 _ => {}
432 }
433 }
434 Ok(())
435}
436
437fn subscription_response(server: &McpServer, req: &Value) -> Response {
443 let id = req.get("id").cloned().unwrap_or(Value::Null);
444 let (ack, close) = server.subscription_ack_and_close(&id);
445 let events = vec![
446 Ok::<_, Infallible>(Event::default().data(ack.to_string())),
447 Ok(Event::default().data(rpc_result(id, close).to_string())),
448 ];
449 (
450 [(HeaderName::from_static("x-accel-buffering"), "no")],
451 Sse::new(tokio_stream::iter(events)),
452 )
453 .into_response()
454}
455
456async fn rpc_handler(
457 State(state): State<HttpTransport>,
458 headers: HeaderMap,
459 Json(req): Json<Value>,
460) -> Response {
461 let id = req.get("id").cloned();
462 if !origin_allowed(&headers) {
463 return (StatusCode::FORBIDDEN, forbidden_response(id)).into_response();
464 }
465 if !authorized(&state.server, &headers) {
466 return (
469 StatusCode::UNAUTHORIZED,
470 [(header::WWW_AUTHENTICATE, r#"Bearer realm="mcp""#)],
471 unauthorized_response(id),
472 )
473 .into_response();
474 }
475
476 if let Some(batch) = req.as_array()
485 && batch.iter().any(|r| request_protocol_version(r).is_some())
486 {
487 return (
488 StatusCode::BAD_REQUEST,
489 Json(rpc_error(
490 id,
491 ERR_INVALID_REQUEST,
492 "Invalid Request: JSON-RPC batches are not supported by protocol revisions after 2024-11-05",
493 None,
494 )),
495 )
496 .into_response();
497 }
498 let method = req.get("method").and_then(|v| v.as_str()).unwrap_or("");
499 let modern = if method == "initialize" {
500 None
501 } else {
502 request_protocol_version(&req)
503 };
504 if modern.is_some() && req.get("id").is_some() {
505 if let Err((status, body)) = validate_modern_headers(&req, &headers) {
506 return (status, body).into_response();
507 }
508 if !MODERN_METHODS.contains(&method) {
509 return (
510 StatusCode::NOT_FOUND,
511 Json(rpc_error(
512 id,
513 ERR_METHOD_NOT_FOUND,
514 &format!("Method not found: {method}"),
515 None,
516 )),
517 )
518 .into_response();
519 }
520 if method == "subscriptions/listen" {
521 return subscription_response(&state.server, &req);
522 }
523 }
524
525 match dispatch_any(&state.server, &req).await {
526 Some(resp) => (StatusCode::OK, Json(resp)).into_response(),
527 None => StatusCode::ACCEPTED.into_response(),
530 }
531}
532
533#[cfg(test)]
534mod tests {
535 use super::*;
536 use crate::mcp::schema::{ResourceDescription, ToolDescription};
537
538 fn server_with_echo() -> McpServer {
539 let mut server = McpServer::new("http-test", "1.0.0");
540 server.register_tool(ToolDescription {
541 name: "echo".into(),
542 description: "Echo".into(),
543 input_schema: serde_json::json!({"type": "object"}),
544 });
545 server.set_async_handler("echo", |params| async move { Ok(params) });
546 server
547 }
548
549 #[tokio::test]
550 async fn dispatch_initialize() {
551 let server = server_with_echo();
552 let req = serde_json::json!({"jsonrpc":"2.0","id":1,"method":"initialize","params":{}});
553 let resp = dispatch_async(&server, &req).await.unwrap();
554 assert_eq!(resp["result"]["serverInfo"]["name"], "http-test");
555 }
556
557 #[tokio::test]
558 async fn dispatch_tools_call_async() {
559 let server = server_with_echo();
560 let req = serde_json::json!({
561 "jsonrpc": "2.0", "id": 2, "method": "tools/call",
562 "params": { "name": "echo", "arguments": { "msg": "hello" } }
563 });
564 let resp = dispatch_async(&server, &req).await.unwrap();
565 let text = resp["result"]["content"][0]["text"].as_str().unwrap();
566 assert!(text.contains("hello"));
567 }
568
569 #[tokio::test]
570 async fn dispatch_unknown_method() {
571 let server = server_with_echo();
572 let req = serde_json::json!({"jsonrpc":"2.0","id":3,"method":"nope"});
573 let resp = dispatch_async(&server, &req).await.unwrap();
574 assert_eq!(resp["error"]["code"], ERR_METHOD_NOT_FOUND);
575 }
576
577 #[tokio::test]
579 async fn dispatch_resources_read() {
580 let mut server = McpServer::new("http-test", "1.0.0");
581 server.register_resource(ResourceDescription {
582 uri: "docs://x".into(),
583 name: "X".into(),
584 description: None,
585 mime_type: None,
586 });
587 server.set_resource_handler("docs://x", |_| Ok(serde_json::json!("# body")));
588 let req = serde_json::json!({
589 "jsonrpc": "2.0", "id": 4, "method": "resources/read",
590 "params": { "uri": "docs://x" }
591 });
592 let resp = dispatch_async(&server, &req).await.unwrap();
593 let text = resp["result"]["contents"][0]["text"].as_str().unwrap();
594 assert!(text.contains("body"));
595 }
596
597 #[test]
598 fn origin_validation_blocks_cross_site_browsers() {
599 let mut h = HeaderMap::new();
600 assert!(origin_allowed(&h), "no Origin (non-browser client) passes");
601 h.insert("origin", "http://localhost:3000".parse().unwrap());
602 assert!(origin_allowed(&h));
603 h.insert("origin", "http://127.0.0.1:8080".parse().unwrap());
604 assert!(origin_allowed(&h));
605 h.insert("origin", "https://evil.example.com".parse().unwrap());
606 assert!(!origin_allowed(&h), "DNS-rebinding origin must be rejected");
607 h.insert("origin", "null".parse().unwrap());
608 assert!(!origin_allowed(&h));
609 h.insert("origin", "https://localhost.evil.com".parse().unwrap());
611 assert!(!origin_allowed(&h));
612 h.insert("origin", "http://[::1]:3000".parse().unwrap());
614 assert!(origin_allowed(&h), "IPv6 loopback must pass: {h:?}");
615 h.insert("origin", "http://[::1]".parse().unwrap());
616 assert!(origin_allowed(&h), "IPv6 loopback (no port) must pass");
617 h.insert("origin", "http://[fe80::1]:3000".parse().unwrap());
618 assert!(!origin_allowed(&h), "non-loopback IPv6 must be rejected");
619 }
620
621 #[tokio::test]
622 async fn batch_requests_get_a_batch_response() {
623 let server = server_with_echo();
624 let batch = serde_json::json!([
625 {"jsonrpc":"2.0","id":1,"method":"ping"},
626 {"jsonrpc":"2.0","id":2,"method":"tools/call",
627 "params":{"name":"echo","arguments":{"v":1}}}
628 ]);
629 let resp = dispatch_any(&server, &batch)
630 .await
631 .expect("batch must not be dropped");
632 let arr = resp.as_array().expect("array response");
633 assert_eq!(arr.len(), 2);
634 assert_eq!(arr[0]["id"], 1);
635 assert_eq!(arr[1]["result"]["isError"], false);
636 }
637
638 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
641 async fn http_round_trip_calls_tool() {
642 let server = Arc::new(server_with_echo());
643 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
644 let addr = listener.local_addr().unwrap();
645 let transport = HttpTransport::new(server);
647 tokio::spawn(async move {
648 let _ = axum::serve(listener, transport.router()).await;
649 });
650
651 let body = serde_json::to_string(&serde_json::json!({
652 "jsonrpc": "2.0", "id": 9, "method": "tools/call",
653 "params": { "name": "echo", "arguments": { "v": 42 } }
654 }))
655 .unwrap();
656 let req = format!(
657 "POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
658 body.len(),
659 body
660 );
661
662 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
663 use tokio::io::{AsyncReadExt, AsyncWriteExt};
664 stream.write_all(req.as_bytes()).await.unwrap();
665 let mut buf = Vec::new();
666 stream.read_to_end(&mut buf).await.unwrap();
667 let response = String::from_utf8_lossy(&buf);
668 assert!(response.contains("200 OK"), "response: {response}");
669 assert!(response.contains("\"content\""), "response: {response}");
672 assert!(response.contains("\\\"v\\\":42"), "response: {response}");
673 }
674
675 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
678 async fn unauthorized_response_carries_www_authenticate() {
679 let server = Arc::new(McpServer::new("http-test", "1.0.0").with_bearer_auth("tok"));
680 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
681 let addr = listener.local_addr().unwrap();
682 let transport = HttpTransport::new(server);
683 tokio::spawn(async move {
684 let _ = axum::serve(listener, transport.router()).await;
685 });
686
687 let body = r#"{"jsonrpc":"2.0","id":1,"method":"tools/list"}"#;
688 let req = format!(
689 "POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
690 body.len(),
691 body
692 );
693 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
694 use tokio::io::{AsyncReadExt, AsyncWriteExt};
695 stream.write_all(req.as_bytes()).await.unwrap();
696 let mut buf = Vec::new();
697 stream.read_to_end(&mut buf).await.unwrap();
698 let response = String::from_utf8_lossy(&buf);
699 assert!(
700 response.contains("401 Unauthorized"),
701 "response: {response}"
702 );
703 assert!(
704 response
705 .to_ascii_lowercase()
706 .contains("www-authenticate: bearer"),
707 "missing WWW-Authenticate challenge: {response}"
708 );
709 }
710
711 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
714 async fn notification_only_post_returns_202() {
715 let server = Arc::new(server_with_echo());
716 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
717 let addr = listener.local_addr().unwrap();
718 let transport = HttpTransport::new(server);
719 tokio::spawn(async move {
720 let _ = axum::serve(listener, transport.router()).await;
721 });
722
723 let body = r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#;
724 let req = format!(
725 "POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
726 body.len(),
727 body
728 );
729 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
730 use tokio::io::{AsyncReadExt, AsyncWriteExt};
731 stream.write_all(req.as_bytes()).await.unwrap();
732 let mut buf = Vec::new();
733 stream.read_to_end(&mut buf).await.unwrap();
734 let response = String::from_utf8_lossy(&buf);
735 assert!(response.contains("202 Accepted"), "response: {response}");
736 let body = response.split("\r\n\r\n").nth(1).unwrap_or("");
737 assert!(body.trim().is_empty(), "202 must have no body: {body}");
738 }
739
740 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
743 async fn legacy_sse_endpoint_is_gone() {
744 let server = Arc::new(server_with_echo());
745 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
746 let addr = listener.local_addr().unwrap();
747 let transport = HttpTransport::new(server);
748 tokio::spawn(async move {
749 let _ = axum::serve(listener, transport.router()).await;
750 });
751
752 let body = r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#;
753 let req = format!(
754 "POST /mcp/sse HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
755 body.len(),
756 body
757 );
758 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
759 use tokio::io::{AsyncReadExt, AsyncWriteExt};
760 stream.write_all(req.as_bytes()).await.unwrap();
761 let mut buf = Vec::new();
762 stream.read_to_end(&mut buf).await.unwrap();
763 let response = String::from_utf8_lossy(&buf);
764 assert!(response.contains("404 Not Found"), "response: {response}");
765 }
766
767 async fn post_raw(headers: &[(&str, &str)], body: &str) -> String {
770 let server = Arc::new(server_with_echo());
771 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
772 let addr = listener.local_addr().unwrap();
773 let transport = HttpTransport::new(server);
774 tokio::spawn(async move {
775 let _ = axum::serve(listener, transport.router()).await;
776 });
777
778 let header_block: String = headers
779 .iter()
780 .map(|(k, v)| format!("{k}: {v}\r\n"))
781 .collect();
782 let req = format!(
783 "POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: {}\r\n{header_block}Connection: close\r\n\r\n{}",
784 body.len(),
785 body
786 );
787 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
788 use tokio::io::{AsyncReadExt, AsyncWriteExt};
789 stream.write_all(req.as_bytes()).await.unwrap();
790 let mut buf = Vec::new();
791 stream.read_to_end(&mut buf).await.unwrap();
792 String::from_utf8_lossy(&buf).into_owned()
793 }
794
795 fn modern_body(method: &str, id: i64) -> String {
797 serde_json::json!({
798 "jsonrpc": "2.0", "id": id, "method": method,
799 "params": { "_meta": { "io.modelcontextprotocol/protocolVersion": "2026-07-28" } }
800 })
801 .to_string()
802 }
803
804 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
807 async fn modern_request_without_headers_is_400_header_mismatch() {
808 let response = post_raw(&[], &modern_body("tools/list", 1)).await;
809 assert!(response.contains("400 Bad Request"), "response: {response}");
810 assert!(response.contains("-32020"), "response: {response}");
811 }
812
813 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
816 async fn modern_request_with_valid_headers_answers_shaped_result() {
817 let response = post_raw(
818 &[
819 ("MCP-Protocol-Version", "2026-07-28"),
820 ("Mcp-Method", "tools/list"),
821 ],
822 &modern_body("tools/list", 2),
823 )
824 .await;
825 assert!(response.contains("200 OK"), "response: {response}");
826 assert!(
827 response.contains("\"resultType\":\"complete\""),
828 "response: {response}"
829 );
830 assert!(response.contains("\"ttlMs\""), "response: {response}");
831 assert!(
832 response.contains("\"cacheScope\":\"private\""),
833 "response: {response}"
834 );
835 }
836
837 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
839 async fn modern_version_header_mismatch_is_400() {
840 let response = post_raw(
841 &[
842 ("MCP-Protocol-Version", "2025-06-18"),
843 ("Mcp-Method", "tools/list"),
844 ],
845 &modern_body("tools/list", 3),
846 )
847 .await;
848 assert!(response.contains("400 Bad Request"), "response: {response}");
849 assert!(response.contains("-32020"), "response: {response}");
850 }
851
852 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
855 async fn modern_unsupported_version_is_400_with_supported_list() {
856 let body = serde_json::json!({
857 "jsonrpc": "2.0", "id": 4, "method": "tools/list",
858 "params": { "_meta": { "io.modelcontextprotocol/protocolVersion": "1900-01-01" } }
859 })
860 .to_string();
861 let response = post_raw(
862 &[
863 ("MCP-Protocol-Version", "1900-01-01"),
864 ("Mcp-Method", "tools/list"),
865 ],
866 &body,
867 )
868 .await;
869 assert!(response.contains("400 Bad Request"), "response: {response}");
870 assert!(response.contains("-32022"), "response: {response}");
871 assert!(response.contains("\"supported\""), "response: {response}");
872 }
873
874 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
877 async fn modern_unknown_method_is_404() {
878 let response = post_raw(
879 &[
880 ("MCP-Protocol-Version", "2026-07-28"),
881 ("Mcp-Method", "no/such-method"),
882 ],
883 &modern_body("no/such-method", 5),
884 )
885 .await;
886 assert!(response.contains("404 Not Found"), "response: {response}");
887 assert!(response.contains("-32601"), "response: {response}");
888 }
889
890 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
892 async fn modern_ping_is_404() {
893 let response = post_raw(
894 &[
895 ("MCP-Protocol-Version", "2026-07-28"),
896 ("Mcp-Method", "ping"),
897 ],
898 &modern_body("ping", 6),
899 )
900 .await;
901 assert!(response.contains("404 Not Found"), "response: {response}");
902 }
903
904 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
907 async fn modern_request_inside_batch_rejects_the_batch() {
908 let batch = format!("[{}]", modern_body("tools/list", 9));
909 let response = post_raw(&[], &batch).await;
910 assert!(response.contains("400 Bad Request"), "response: {response}");
911 assert!(response.contains("-32600"), "response: {response}");
912 }
913
914 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
916 async fn legacy_batch_is_served() {
917 let batch = r#"[{"jsonrpc":"2.0","id":1,"method":"tools/list"},{"jsonrpc":"2.0","id":2,"method":"ping"}]"#;
918 let response = post_raw(&[], batch).await;
919 assert!(response.contains("200 OK"), "response: {response}");
920 assert!(response.contains("\"tools\""), "response: {response}");
921 }
922
923 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
925 async fn modern_tools_call_requires_matching_mcp_name() {
926 let body = serde_json::json!({
927 "jsonrpc": "2.0", "id": 7, "method": "tools/call",
928 "params": {
929 "name": "echo",
930 "arguments": { "v": 1 },
931 "_meta": { "io.modelcontextprotocol/protocolVersion": "2026-07-28" }
932 }
933 })
934 .to_string();
935 let response = post_raw(
937 &[
938 ("MCP-Protocol-Version", "2026-07-28"),
939 ("Mcp-Method", "tools/call"),
940 ],
941 &body,
942 )
943 .await;
944 assert!(response.contains("400 Bad Request"), "response: {response}");
945 let response = post_raw(
947 &[
948 ("MCP-Protocol-Version", "2026-07-28"),
949 ("Mcp-Method", "tools/call"),
950 ("Mcp-Name", "echo"),
951 ],
952 &body,
953 )
954 .await;
955 assert!(response.contains("200 OK"), "response: {response}");
956 assert!(
957 response.contains("\"resultType\":\"complete\""),
958 "response: {response}"
959 );
960 }
961
962 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
966 async fn modern_subscriptions_listen_streams_ack_and_close() {
967 let response = post_raw(
968 &[
969 ("MCP-Protocol-Version", "2026-07-28"),
970 ("Mcp-Method", "subscriptions/listen"),
971 ],
972 &modern_body("subscriptions/listen", 8),
973 )
974 .await;
975 assert!(response.contains("200 OK"), "response: {response}");
976 assert!(
977 response
978 .to_ascii_lowercase()
979 .contains("content-type: text/event-stream"),
980 "response: {response}"
981 );
982 assert!(
983 response
984 .to_ascii_lowercase()
985 .contains("x-accel-buffering: no"),
986 "response: {response}"
987 );
988 assert!(
989 response.contains("notifications/subscriptions/acknowledged"),
990 "response: {response}"
991 );
992 assert!(
993 response.contains("\"resultType\":\"complete\""),
994 "closure result on the stream: {response}"
995 );
996 }
997
998 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1001 async fn legacy_request_needs_no_modern_headers() {
1002 let body = r#"{"jsonrpc":"2.0","id":9,"method":"tools/list"}"#;
1003 let response = post_raw(&[], body).await;
1004 assert!(response.contains("200 OK"), "response: {response}");
1005 assert!(
1007 !response.contains("resultType"),
1008 "legacy result must stay unshaped: {response}"
1009 );
1010 }
1011
1012 #[test]
1014 fn decode_header_value_handles_base64_sentinel() {
1015 assert_eq!(decode_header_value("=?base64?ZWNobw==?="), "echo");
1017 assert_eq!(decode_header_value("echo"), "echo");
1018 assert_eq!(decode_header_value("=?base64?!!!?="), "=?base64?!!!?=");
1020 }
1021
1022 #[tokio::test]
1024 async fn dispatch_modern_discover() {
1025 let server = server_with_echo();
1026 let req: Value = serde_json::from_str(&modern_body("server/discover", 1)).unwrap();
1027 let resp = dispatch_async(&server, &req).await.unwrap();
1028 assert_eq!(
1029 resp["result"]["supportedVersions"][0],
1030 crate::mcp::server::LATEST_PROTOCOL_VERSION
1031 );
1032 assert_eq!(resp["result"]["resultType"], "complete");
1033 assert!(resp["result"]["ttlMs"].is_u64());
1034 }
1035}