1use async_trait::async_trait;
17use axum::{
18 extract::State,
19 http::{HeaderMap, StatusCode},
20 response::{IntoResponse, Response},
21 routing::{get, post},
22 Json, Router,
23};
24
25use axum::response::{sse::Event, Sse};
26use reqwest::Client;
27use serde_json::Value;
28use std::{collections::HashMap, convert::Infallible, sync::Arc, time::Duration};
29use tokio::sync::{broadcast, mpsc, Mutex, RwLock};
30
31#[cfg(feature = "sse")]
32use futures::Stream;
33use futures::StreamExt;
34
35#[cfg(feature = "sse")]
36use tokio_stream::wrappers::BroadcastStream;
37
38use tower::ServiceBuilder;
39use tower_http::cors::{Any, CorsLayer};
40use tracing::Instrument;
41
42#[cfg(feature = "tls")]
43use hyper_util::{
44 rt::{TokioExecutor, TokioIo},
45 server::conn::auto::Builder as HyperServerBuilder,
46 service::TowerToHyperService,
47};
48
49use crate::core::error::{McpError, McpResult};
50use crate::core::logging::ErrorContext;
51use crate::protocol::{
52 encode_http_header_value, has_tasks_extension, json_rpc_error_details, methods,
53 modern_request_context, request_protocol_version, request_routing_name, tool_call_headers,
54 tool_header_mappings,
55 types::{
56 error_codes, JsonRpcError, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest,
57 JsonRpcResponse,
58 },
59 validate_http_headers, validate_tool_call_headers, ServerCapabilities, SubscriptionFilter,
60 SubscriptionsAcknowledgedParams, SubscriptionsListenParams, HEADER_MISMATCH, MCP_METHOD_HEADER,
61 MCP_NAME_HEADER, MCP_PROTOCOL_VERSION_HEADER, SUBSCRIPTION_ID_META_KEY, TASKS_EXTENSION_ID,
62 UNSUPPORTED_PROTOCOL_VERSION,
63};
64use crate::transport::traits::{
65 ClientSubscription, ConnectionState, ServerTransport, Transport, TransportConfig,
66};
67
68const FORBIDDEN_ERROR: i32 = -32010;
69const RATE_LIMITED_ERROR: i32 = -32011;
70
71fn parse_sse_response(bytes: &[u8], request_id: &Value) -> McpResult<Value> {
72 let body = String::from_utf8_lossy(bytes).replace("\r\n", "\n");
73 for event in body.split("\n\n") {
74 let data = event
75 .lines()
76 .filter_map(|line| line.strip_prefix("data:"))
77 .map(str::trim_start)
78 .collect::<Vec<_>>()
79 .join("\n");
80 if data.is_empty() {
81 continue;
82 }
83 let value: Value = serde_json::from_str(&data)
84 .map_err(|error| McpError::Serialization(format!("invalid SSE JSON data: {error}")))?;
85 if value.get("id") == Some(request_id)
86 && (value.get("result").is_some() || value.get("error").is_some())
87 {
88 return Ok(value);
89 }
90 }
91 Err(McpError::Serialization(
92 "SSE response ended without a JSON-RPC result for the request".to_string(),
93 ))
94}
95
96#[cfg(feature = "otel")]
97struct HeaderExtractor<'a>(&'a HeaderMap);
98
99#[cfg(feature = "otel")]
100impl opentelemetry::propagation::Extractor for HeaderExtractor<'_> {
101 fn get(&self, key: &str) -> Option<&str> {
102 self.0.get(key).and_then(|value| value.to_str().ok())
103 }
104
105 fn keys(&self) -> Vec<&str> {
106 self.0.keys().map(axum::http::HeaderName::as_str).collect()
107 }
108}
109
110#[cfg(feature = "otel")]
111struct MapInjector<'a>(&'a mut HashMap<String, String>);
112
113#[cfg(feature = "otel")]
114impl opentelemetry::propagation::Injector for MapInjector<'_> {
115 fn set(&mut self, key: &str, value: String) {
116 self.0.insert(key.to_string(), value);
117 }
118}
119
120#[cfg(feature = "otel")]
121fn inject_trace_context(mut request: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
122 use tracing_opentelemetry::OpenTelemetrySpanExt;
123
124 let context = tracing::Span::current().context();
125 let mut headers = HashMap::new();
126 opentelemetry::global::get_text_map_propagator(|propagator| {
127 propagator.inject_context(&context, &mut MapInjector(&mut headers));
128 });
129 for (key, value) in headers {
130 request = request.header(key, value);
131 }
132 request
133}
134
135#[cfg(feature = "tls")]
137#[derive(Debug, Clone)]
138pub struct MtlsClientConfig {
139 pub identity_pem: Vec<u8>,
140 pub ca_certificate_pem: Vec<u8>,
141}
142
143#[cfg(feature = "tls")]
144impl MtlsClientConfig {
145 pub fn new(identity_pem: impl Into<Vec<u8>>, ca_certificate_pem: impl Into<Vec<u8>>) -> Self {
146 Self {
147 identity_pem: identity_pem.into(),
148 ca_certificate_pem: ca_certificate_pem.into(),
149 }
150 }
151}
152
153#[cfg(feature = "tls")]
155#[derive(Debug, Clone)]
156pub struct MtlsServerConfig {
157 pub certificate_chain_pem: Vec<u8>,
158 pub private_key_pem: Vec<u8>,
159 pub client_ca_pem: Vec<u8>,
160}
161
162#[cfg(feature = "tls")]
163impl MtlsServerConfig {
164 pub fn new(
165 certificate_chain_pem: impl Into<Vec<u8>>,
166 private_key_pem: impl Into<Vec<u8>>,
167 client_ca_pem: impl Into<Vec<u8>>,
168 ) -> Self {
169 Self {
170 certificate_chain_pem: certificate_chain_pem.into(),
171 private_key_pem: private_key_pem.into(),
172 client_ca_pem: client_ca_pem.into(),
173 }
174 }
175
176 fn build_rustls(&self) -> McpResult<rustls::ServerConfig> {
177 use rustls::pki_types::{pem::PemObject, CertificateDer, PrivateKeyDer};
178 use rustls::server::WebPkiClientVerifier;
179 use rustls::RootCertStore;
180
181 let certificates = CertificateDer::pem_slice_iter(&self.certificate_chain_pem)
182 .collect::<Result<Vec<_>, _>>()
183 .map_err(|error| {
184 McpError::Authentication(format!("invalid server certificate: {error}"))
185 })?;
186 if certificates.is_empty() {
187 return Err(McpError::Authentication(
188 "mTLS server certificate chain is empty".to_string(),
189 ));
190 }
191
192 let private_key =
193 PrivateKeyDer::from_pem_slice(&self.private_key_pem).map_err(|error| {
194 McpError::Authentication(format!("invalid server private key: {error}"))
195 })?;
196
197 let client_ca = CertificateDer::pem_slice_iter(&self.client_ca_pem)
198 .collect::<Result<Vec<_>, _>>()
199 .map_err(|error| McpError::Authentication(format!("invalid client CA: {error}")))?;
200 let mut roots = RootCertStore::empty();
201 let (accepted, rejected) = roots.add_parsable_certificates(client_ca);
202 if accepted == 0 || rejected > 0 {
203 return Err(McpError::Authentication(format!(
204 "client CA contained {accepted} accepted and {rejected} rejected certificates"
205 )));
206 }
207
208 let verifier = WebPkiClientVerifier::builder(Arc::new(roots))
209 .build()
210 .map_err(|error| {
211 McpError::Authentication(format!("invalid client verifier: {error}"))
212 })?;
213 rustls::ServerConfig::builder_with_protocol_versions(&[&rustls::version::TLS13])
214 .with_client_cert_verifier(verifier)
215 .with_single_cert(certificates, private_key)
216 .map_err(|error| McpError::Authentication(format!("invalid server identity: {error}")))
217 }
218}
219
220#[derive(Debug)]
229pub struct HttpClientTransport {
230 pub(crate) client: Client,
231 pub(crate) base_url: String,
232 pub(crate) sse_url: Option<String>,
233 pub(crate) headers: HeaderMap,
234 pending_requests: Arc<Mutex<HashMap<Value, tokio::sync::oneshot::Sender<JsonRpcResponse>>>>,
236 notification_receiver: Option<mpsc::UnboundedReceiver<JsonRpcNotification>>,
237 pub(crate) config: TransportConfig,
238 state: ConnectionState,
239 request_id_counter: Arc<Mutex<u64>>,
240 tool_schemas: HashMap<String, Value>,
242 subscription_tasks: Arc<Mutex<HashMap<String, tokio::task::AbortHandle>>>,
243}
244
245impl HttpClientTransport {
246 pub async fn new<S: AsRef<str>>(base_url: S, sse_url: Option<S>) -> McpResult<Self> {
255 Self::with_config(base_url, sse_url, TransportConfig::default()).await
256 }
257
258 pub async fn with_config<S: AsRef<str>>(
268 base_url: S,
269 sse_url: Option<S>,
270 config: TransportConfig,
271 ) -> McpResult<Self> {
272 let client_builder = Client::builder()
273 .timeout(Duration::from_millis(
274 config.read_timeout_ms.unwrap_or(60_000),
275 ))
276 .connect_timeout(Duration::from_millis(
277 config.connect_timeout_ms.unwrap_or(30_000),
278 ));
279
280 let client = client_builder
283 .build()
284 .map_err(|e| McpError::Http(format!("Failed to create HTTP client: {e}")))?;
285
286 let mut headers = HeaderMap::new();
287 headers.insert("Content-Type", "application/json".parse().unwrap());
288 headers.insert(
289 "Accept",
290 "application/json, text/event-stream".parse().unwrap(),
291 );
292
293 for (key, value) in &config.headers {
295 if let (Ok(header_name), Ok(header_value)) = (
296 key.parse::<axum::http::HeaderName>(),
297 value.parse::<axum::http::HeaderValue>(),
298 ) {
299 headers.insert(header_name, header_value);
300 }
301 }
302
303 let (notification_sender, notification_receiver) = mpsc::unbounded_channel();
304
305 if let Some(sse_url) = &sse_url {
307 let sse_url = sse_url.as_ref().to_string();
308 let client_clone = client.clone();
309 let headers_clone = headers.clone();
310
311 tokio::spawn(async move {
312 if let Err(e) = Self::handle_sse_stream(
313 client_clone,
314 sse_url,
315 headers_clone,
316 notification_sender,
317 )
318 .await
319 {
320 tracing::error!("SSE stream error: {}", e);
321 }
322 });
323 }
324
325 Ok(Self {
326 client,
327 base_url: base_url.as_ref().to_string(),
328 sse_url: sse_url.map(|s| s.as_ref().to_string()),
329 headers,
330 pending_requests: Arc::new(Mutex::new(HashMap::new())),
331 notification_receiver: Some(notification_receiver),
332 config,
333 state: ConnectionState::Connected,
334 request_id_counter: Arc::new(Mutex::new(0)),
335 tool_schemas: HashMap::new(),
336 subscription_tasks: Arc::new(Mutex::new(HashMap::new())),
337 })
338 }
339
340 #[cfg(feature = "tls")]
343 pub async fn with_mtls<S: AsRef<str>>(
344 base_url: S,
345 sse_url: Option<S>,
346 config: TransportConfig,
347 mtls: MtlsClientConfig,
348 ) -> McpResult<Self> {
349 let identity = reqwest::Identity::from_pem(&mtls.identity_pem).map_err(|error| {
350 McpError::Authentication(format!("invalid client identity: {error}"))
351 })?;
352 let root = reqwest::Certificate::from_pem(&mtls.ca_certificate_pem)
353 .map_err(|error| McpError::Authentication(format!("invalid server CA: {error}")))?;
354 let client = Client::builder()
355 .timeout(Duration::from_millis(
356 config.read_timeout_ms.unwrap_or(60_000),
357 ))
358 .connect_timeout(Duration::from_millis(
359 config.connect_timeout_ms.unwrap_or(30_000),
360 ))
361 .identity(identity)
362 .tls_certs_only([root])
363 .min_tls_version(reqwest::tls::Version::TLS_1_3)
364 .build()
365 .map_err(|error| McpError::Http(format!("failed to create mTLS client: {error}")))?;
366
367 let base = base_url.as_ref().to_string();
368 let sse = sse_url.as_ref().map(|url| url.as_ref().to_string());
369 let mut transport = Self::with_config(base.as_str(), None::<&str>, config).await?;
370 transport.client = client.clone();
371 transport.sse_url = sse.clone();
372
373 if let Some(url) = sse {
374 let headers = transport.headers.clone();
375 let sender = {
376 let (sender, receiver) = mpsc::unbounded_channel();
377 transport.notification_receiver = Some(receiver);
378 sender
379 };
380 tokio::spawn(async move {
381 if let Err(error) = Self::handle_sse_stream(client, url, headers, sender).await {
382 tracing::error!(%error, "mTLS SSE stream failed");
383 }
384 });
385 }
386 Ok(transport)
387 }
388
389 async fn handle_sse_stream(
390 client: Client,
391 sse_url: String,
392 headers: HeaderMap,
393 notification_sender: mpsc::UnboundedSender<JsonRpcNotification>,
394 ) -> McpResult<()> {
395 let mut request = client.get(&sse_url);
396 #[cfg(feature = "otel")]
397 {
398 request = inject_trace_context(request);
399 }
400 for (name, value) in headers.iter() {
401 let name_str = name.as_str();
403 let value_bytes = value.as_bytes();
404 request = request.header(name_str, value_bytes);
405 }
406
407 let _response = request
408 .send()
409 .await
410 .map_err(|e| McpError::Http(format!("SSE connection failed: {e}")))?;
411
412 #[cfg(feature = "sse")]
413 {
414 let mut stream = _response.bytes_stream();
415 while let Some(chunk) = stream.next().await {
416 match chunk {
417 Ok(bytes) => {
418 let text = String::from_utf8_lossy(&bytes);
419 for line in text.lines() {
420 if let Some(data) = line.strip_prefix("data: ") {
421 if let Ok(notification) =
423 serde_json::from_str::<JsonRpcNotification>(data)
424 {
425 if notification_sender.send(notification).is_err() {
426 tracing::debug!("Notification receiver dropped");
427 return Ok(());
428 }
429 }
430 }
431 }
432 }
433 Err(e) => {
434 tracing::error!("SSE stream error: {}", e);
435 break;
436 }
437 }
438 }
439 }
440
441 #[cfg(not(feature = "sse"))]
442 {
443 let _ = notification_sender; tracing::warn!("SSE streaming requires SSE feature");
445 }
446
447 Ok(())
448 }
449
450 pub async fn next_request_id(&self) -> u64 {
451 let mut counter = self.request_id_counter.lock().await;
452 *counter += 1;
453 *counter
454 }
455
456 fn mcp_url(&self) -> String {
457 let base = self.base_url.trim_end_matches('/');
458 if base.ends_with("/mcp") {
459 base.to_string()
460 } else {
461 format!("{base}/mcp")
462 }
463 }
464
465 async fn track_request(&self, request_id: &Value) {
467 let mut pending = self.pending_requests.lock().await;
471 let (sender, _receiver) = tokio::sync::oneshot::channel();
472 pending.insert(request_id.clone(), sender);
473 }
474
475 async fn untrack_request(&self, request_id: &Value) {
477 let mut pending = self.pending_requests.lock().await;
478 pending.remove(request_id);
479 }
480
481 pub async fn active_request_count(&self) -> usize {
483 let pending = self.pending_requests.lock().await;
484 pending.len()
485 }
486
487 fn capture_tool_schemas(&mut self, response: &mut JsonRpcResponse) {
488 let Some(tools) = response
489 .result
490 .as_mut()
491 .and_then(Value::as_object_mut)
492 .and_then(|result| result.get_mut("tools"))
493 .and_then(Value::as_array_mut)
494 else {
495 return;
496 };
497 self.tool_schemas.clear();
498 tools.retain(|tool| {
499 let Some(name) = tool.get("name").and_then(Value::as_str) else {
500 return false;
501 };
502 let Some(schema) = tool.get("inputSchema") else {
503 return false;
504 };
505 match tool_header_mappings(schema) {
506 Ok(_) => {
507 self.tool_schemas.insert(name.to_string(), schema.clone());
508 true
509 }
510 Err(error) => {
511 tracing::warn!(tool.name = name, %error, "excluding tool with invalid x-mcp-header schema");
512 false
513 }
514 }
515 });
516 }
517
518 #[cfg(test)]
519 pub fn has_notification_receiver(&self) -> bool {
520 self.notification_receiver.is_some()
521 }
522}
523
524#[async_trait]
525impl Transport for HttpClientTransport {
526 async fn send_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
527 let request_with_id = if request.id == Value::Null {
529 let request_id = self.next_request_id().await;
530 JsonRpcRequest {
531 id: Value::from(request_id),
532 ..request
533 }
534 } else {
535 request
536 };
537
538 let context = ErrorContext::new("http_send_request")
540 .with_transport("http")
541 .with_method(&request_with_id.method)
542 .with_extra("request_id", request_with_id.id.clone())
543 .with_extra("base_url", serde_json::Value::String(self.base_url.clone()));
544
545 self.track_request(&request_with_id.id).await;
547
548 let url = self.mcp_url();
549
550 let mut http_request = self.client.post(&url);
551
552 #[cfg(feature = "otel")]
553 {
554 http_request = inject_trace_context(http_request);
555 }
556
557 for (name, value) in self.headers.iter() {
559 let name_str = name.as_str();
560 let value_bytes = value.as_bytes();
561 http_request = http_request.header(name_str, value_bytes);
562 }
563
564 if let Some(version) = request_protocol_version(&request_with_id) {
565 http_request = http_request
566 .header(MCP_PROTOCOL_VERSION_HEADER, version)
567 .header(MCP_METHOD_HEADER, request_with_id.method.as_str());
568 if let Some(name) = request_routing_name(&request_with_id) {
569 http_request = http_request.header(MCP_NAME_HEADER, encode_http_header_value(name));
570 }
571 if request_with_id.method == methods::TOOLS_CALL {
572 if let Some(params) = request_with_id.params.as_ref().and_then(Value::as_object) {
573 if let Some(tool_name) = params.get("name").and_then(Value::as_str) {
574 if let Some(schema) = self.tool_schemas.get(tool_name) {
575 let arguments = params.get("arguments").unwrap_or(&Value::Null);
576 for (name, value) in tool_call_headers(schema, arguments)? {
577 http_request = http_request.header(name, value);
578 }
579 }
580 }
581 }
582 }
583 }
584
585 if let Some(timeout_ms) = self.config.read_timeout_ms {
587 http_request = http_request.timeout(Duration::from_millis(timeout_ms));
588 }
589
590 let response = http_request
591 .json(&request_with_id)
592 .send()
593 .await
594 .map_err(|e| {
595 let request_id = request_with_id.id.clone();
597 let pending_requests = self.pending_requests.clone();
598 tokio::spawn(async move {
599 let mut pending = pending_requests.lock().await;
600 pending.remove(&request_id);
601 });
602
603 let error = if e.is_timeout() {
605 McpError::timeout("HTTP request timeout")
606 } else if e.is_connect() {
607 McpError::connection(format!("HTTP connection failed: {e}"))
608 } else {
609 McpError::Http(format!("HTTP request failed: {e}"))
610 };
611
612 let error_clone = error.clone();
614 let context_clone = context.clone();
615 tokio::spawn(async move {
616 error_clone.log_with_context(context_clone).await;
617 });
618
619 error
620 })?;
621
622 let response_status = response.status();
623 let response_content_type = response
624 .headers()
625 .get(reqwest::header::CONTENT_TYPE)
626 .and_then(|value| value.to_str().ok())
627 .unwrap_or_default()
628 .to_string();
629 let response_bytes = response
630 .bytes()
631 .await
632 .map_err(|error| McpError::Http(format!("failed to read HTTP response: {error}")))?;
633 let parsed_response = if response_content_type.starts_with("text/event-stream") {
634 parse_sse_response(&response_bytes, &request_with_id.id)
635 } else {
636 serde_json::from_slice(&response_bytes)
637 .map_err(|error| McpError::Serialization(format!("invalid JSON response: {error}")))
638 };
639 let json_value: Value = match parsed_response {
640 Ok(value) => value,
641 Err(_error) if !response_status.is_success() => {
642 self.untrack_request(&request_with_id.id).await;
643 return Err(McpError::Http(format!(
644 "HTTP error: {} {}",
645 response_status.as_u16(),
646 response_status.canonical_reason().unwrap_or("Unknown")
647 )));
648 }
649 Err(error) => {
650 self.untrack_request(&request_with_id.id).await;
651 error.clone().log_with_context(context).await;
652 return Err(error);
653 }
654 };
655
656 let mut result = if json_value.get("error").is_some() {
657 serde_json::from_value::<JsonRpcError>(json_value)
658 .map_err(|error| McpError::Serialization(error.to_string()))
659 .and_then(|json_error| {
660 if json_error.id != request_with_id.id {
661 Err(McpError::Http(format!(
662 "Error response ID {:?} does not match request ID {:?}",
663 json_error.id, request_with_id.id
664 )))
665 } else {
666 Err(match json_error.error.code {
667 FORBIDDEN_ERROR => McpError::Forbidden(json_error.error.message),
668 RATE_LIMITED_ERROR => McpError::RateLimited {
669 retry_after_ms: json_error
670 .error
671 .data
672 .and_then(|data| data.get("retryAfterMs").cloned())
673 .and_then(|value| value.as_u64())
674 .unwrap_or_default(),
675 },
676 error_codes::METHOD_NOT_FOUND => {
677 McpError::MethodNotFound(json_error.error.message)
678 }
679 HEADER_MISMATCH => McpError::HeaderMismatch(json_error.error.message),
680 crate::protocol::MISSING_REQUIRED_CLIENT_CAPABILITY => {
681 let required = json_error
682 .error
683 .data
684 .and_then(|data| data.get("requiredCapabilities").cloned())
685 .unwrap_or_else(|| serde_json::json!({}));
686 McpError::MissingRequiredClientCapability(required)
687 }
688 UNSUPPORTED_PROTOCOL_VERSION => {
689 let data = json_error.error.data.unwrap_or_default();
690 McpError::UnsupportedProtocolVersion {
691 requested: data
692 .get("requested")
693 .and_then(Value::as_str)
694 .unwrap_or("unknown")
695 .to_string(),
696 supported: data
697 .get("supported")
698 .and_then(Value::as_array)
699 .into_iter()
700 .flatten()
701 .filter_map(Value::as_str)
702 .map(str::to_string)
703 .collect(),
704 }
705 }
706 code => McpError::Protocol(format!(
707 "JSON-RPC error {code}: {}",
708 json_error.error.message
709 )),
710 })
711 }
712 })
713 } else if !response_status.is_success() {
714 Err(McpError::Http(format!(
715 "HTTP error: {} {}",
716 response_status.as_u16(),
717 response_status.canonical_reason().unwrap_or("Unknown")
718 )))
719 } else {
720 serde_json::from_value::<JsonRpcResponse>(json_value)
721 .map_err(|error| McpError::Serialization(error.to_string()))
722 .and_then(|json_response| {
723 if json_response.id != request_with_id.id {
724 Err(McpError::Http(format!(
725 "Response ID {:?} does not match request ID {:?}",
726 json_response.id, request_with_id.id
727 )))
728 } else {
729 Ok(json_response)
730 }
731 })
732 };
733
734 if request_with_id.method == methods::TOOLS_LIST {
735 if let Ok(response) = &mut result {
736 self.capture_tool_schemas(response);
737 }
738 }
739 self.untrack_request(&request_with_id.id).await;
740 result
741 }
742
743 async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
744 let url = self.mcp_url();
745
746 let mut http_request = self.client.post(&url);
747
748 #[cfg(feature = "otel")]
749 {
750 http_request = inject_trace_context(http_request);
751 }
752
753 for (name, value) in self.headers.iter() {
755 let name_str = name.as_str();
756 let value_bytes = value.as_bytes();
757 http_request = http_request.header(name_str, value_bytes);
758 }
759
760 if let Some(timeout_ms) = self.config.write_timeout_ms {
762 http_request = http_request.timeout(Duration::from_millis(timeout_ms));
763 }
764
765 let response = http_request
766 .json(¬ification)
767 .send()
768 .await
769 .map_err(|e| McpError::Http(format!("HTTP notification failed: {e}")))?;
770
771 if !response.status().is_success() {
772 return Err(McpError::Http(format!(
773 "HTTP notification error: {} {}",
774 response.status().as_u16(),
775 response.status().canonical_reason().unwrap_or("Unknown")
776 )));
777 }
778
779 Ok(())
780 }
781
782 async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
783 if let Some(ref mut receiver) = self.notification_receiver {
784 match receiver.try_recv() {
785 Ok(notification) => Ok(Some(notification)),
786 Err(mpsc::error::TryRecvError::Empty) => Ok(None),
787 Err(mpsc::error::TryRecvError::Disconnected) => Err(McpError::Http(
788 "Notification channel disconnected".to_string(),
789 )),
790 }
791 } else {
792 Ok(None)
793 }
794 }
795
796 async fn open_subscription(
797 &mut self,
798 request: JsonRpcRequest,
799 ) -> McpResult<ClientSubscription> {
800 if request.method != methods::SUBSCRIPTIONS_LISTEN {
801 return Err(McpError::InvalidParams(
802 "open_subscription requires subscriptions/listen".to_string(),
803 ));
804 }
805 let version = request_protocol_version(&request).ok_or_else(|| {
806 McpError::InvalidParams("subscription request is missing modern metadata".to_string())
807 })?;
808 let url = self.mcp_url();
809 let mut http_request = self
810 .client
811 .post(url)
812 .header("Content-Type", "application/json")
813 .header("Accept", "application/json, text/event-stream")
814 .header(MCP_PROTOCOL_VERSION_HEADER, version)
815 .header(MCP_METHOD_HEADER, methods::SUBSCRIPTIONS_LISTEN);
816 for (name, value) in self.headers.iter() {
817 if !name.as_str().eq_ignore_ascii_case("accept") {
818 http_request = http_request.header(name.as_str(), value.as_bytes());
819 }
820 }
821 #[cfg(feature = "otel")]
822 {
823 http_request = inject_trace_context(http_request);
824 }
825 let response = http_request
826 .json(&request)
827 .send()
828 .await
829 .map_err(|error| McpError::Http(format!("subscription request failed: {error}")))?;
830 let status = response.status();
831 if !status.is_success() {
832 let bytes = response.bytes().await.unwrap_or_default();
833 if let Ok(error) = serde_json::from_slice::<JsonRpcError>(&bytes) {
834 return Err(match error.error.code {
835 crate::protocol::MISSING_REQUIRED_CLIENT_CAPABILITY => {
836 McpError::MissingRequiredClientCapability(
837 error
838 .error
839 .data
840 .and_then(|value| value.get("requiredCapabilities").cloned())
841 .unwrap_or_else(|| serde_json::json!({})),
842 )
843 }
844 HEADER_MISMATCH => McpError::HeaderMismatch(error.error.message),
845 code => McpError::Protocol(format!(
846 "JSON-RPC error {code}: {}",
847 error.error.message
848 )),
849 });
850 }
851 return Err(McpError::Http(format!(
852 "subscription HTTP error: {}",
853 status.as_u16()
854 )));
855 }
856 let content_type = response
857 .headers()
858 .get(reqwest::header::CONTENT_TYPE)
859 .and_then(|value| value.to_str().ok())
860 .unwrap_or_default();
861 if !content_type.starts_with("text/event-stream") {
862 return Err(McpError::Http(format!(
863 "subscriptions/listen requires text/event-stream, received {content_type}"
864 )));
865 }
866
867 let (notification_tx, notification_rx) = mpsc::unbounded_channel();
868 let (completion_tx, completion_rx) = tokio::sync::oneshot::channel();
869 let request_id = request.id.clone();
870 let key = request_id.to_string();
871 let tasks = self.subscription_tasks.clone();
872 let task_key = key.clone();
873 let task = tokio::spawn(async move {
874 let mut completion_tx = Some(completion_tx);
875 let mut stream = response.bytes_stream();
876 let mut buffer = String::new();
877 while let Some(chunk) = stream.next().await {
878 let chunk = match chunk {
879 Ok(chunk) => chunk,
880 Err(error) => {
881 if let Some(sender) = completion_tx.take() {
882 let _ = sender.send(Err(McpError::Http(format!(
883 "subscription stream failed: {error}"
884 ))));
885 }
886 tasks.lock().await.remove(&task_key);
887 return;
888 }
889 };
890 buffer.push_str(&String::from_utf8_lossy(&chunk));
891 buffer = buffer.replace("\r\n", "\n");
892 while let Some(boundary) = buffer.find("\n\n") {
893 let event = buffer[..boundary].to_string();
894 buffer.drain(..boundary + 2);
895 let data = event
896 .lines()
897 .filter_map(|line| line.strip_prefix("data:"))
898 .map(str::trim_start)
899 .collect::<Vec<_>>()
900 .join("\n");
901 if data.is_empty() {
902 continue;
903 }
904 let Ok(value) = serde_json::from_str::<Value>(&data) else {
905 continue;
906 };
907 if value.get("method").is_some() && value.get("id").is_none() {
908 if let Ok(notification) = serde_json::from_value(value) {
909 if notification_tx.send(notification).is_err() {
910 tasks.lock().await.remove(&task_key);
911 return;
912 }
913 }
914 } else if value.get("result").is_some() {
915 if let Some(sender) = completion_tx.take() {
916 let result = serde_json::from_value(value)
917 .map_err(|error| McpError::Serialization(error.to_string()));
918 let _ = sender.send(result);
919 }
920 tasks.lock().await.remove(&task_key);
921 return;
922 } else if value.get("error").is_some() {
923 if let Some(sender) = completion_tx.take() {
924 let message = value
925 .get("error")
926 .and_then(|error| error.get("message"))
927 .and_then(Value::as_str)
928 .unwrap_or("subscription failed");
929 let _ = sender.send(Err(McpError::Protocol(message.to_string())));
930 }
931 tasks.lock().await.remove(&task_key);
932 return;
933 }
934 }
935 }
936 if let Some(sender) = completion_tx.take() {
937 let _ = sender.send(Err(McpError::Transport(
938 "subscription stream closed without a final response".to_string(),
939 )));
940 }
941 tasks.lock().await.remove(&task_key);
942 });
943 let abort_handle = task.abort_handle();
944 self.subscription_tasks
945 .lock()
946 .await
947 .insert(key, abort_handle.clone());
948 Ok(
949 ClientSubscription::new(request_id, notification_rx, completion_rx)
950 .with_abort_handle(abort_handle),
951 )
952 }
953
954 async fn cancel_subscription(&mut self, request_id: &Value) -> McpResult<()> {
955 if let Some(handle) = self
956 .subscription_tasks
957 .lock()
958 .await
959 .remove(&request_id.to_string())
960 {
961 handle.abort();
962 }
963 Ok(())
964 }
965
966 async fn close(&mut self) -> McpResult<()> {
967 for (_, task) in self.subscription_tasks.lock().await.drain() {
968 task.abort();
969 }
970 self.state = ConnectionState::Disconnected;
971 self.notification_receiver = None;
972 Ok(())
973 }
974
975 fn is_connected(&self) -> bool {
976 matches!(self.state, ConnectionState::Connected)
977 }
978
979 fn connection_info(&self) -> String {
980 format!(
981 "HTTP transport (base: {}, sse: {:?}, state: {:?})",
982 self.base_url, self.sse_url, self.state
983 )
984 }
985}
986
987type HttpRequestHandler = Arc<
992 dyn Fn(JsonRpcRequest) -> tokio::sync::oneshot::Receiver<McpResult<JsonRpcResponse>>
993 + Send
994 + Sync,
995>;
996
997#[derive(Clone)]
999struct HttpServerState {
1000 notification_sender: broadcast::Sender<JsonRpcNotification>,
1001 request_handler: Option<HttpRequestHandler>,
1002 tool_schemas: HashMap<String, Value>,
1003 capabilities: ServerCapabilities,
1004}
1005
1006pub struct HttpServerTransport {
1011 bind_addr: String,
1012 config: TransportConfig,
1013 state: Arc<RwLock<HttpServerState>>,
1014 server_handle: Option<tokio::task::JoinHandle<()>>,
1015 running: Arc<RwLock<bool>>,
1016 pending_request_handler: Option<crate::transport::traits::ServerRequestHandler>,
1017 pending_tool_schemas: HashMap<String, Value>,
1018 pending_capabilities: ServerCapabilities,
1019 pending_task_notifications: Option<broadcast::Receiver<JsonRpcNotification>>,
1020 #[cfg(feature = "tls")]
1021 mtls_config: Option<MtlsServerConfig>,
1022}
1023
1024impl HttpServerTransport {
1025 pub fn new<S: Into<String>>(bind_addr: S) -> Self {
1033 Self::with_config(bind_addr, TransportConfig::default())
1034 }
1035
1036 pub fn with_config<S: Into<String>>(bind_addr: S, config: TransportConfig) -> Self {
1045 let (notification_sender, _) = broadcast::channel(1000);
1046
1047 Self {
1048 bind_addr: bind_addr.into(),
1049 config,
1050 state: Arc::new(RwLock::new(HttpServerState {
1051 notification_sender,
1052 request_handler: None,
1053 tool_schemas: HashMap::new(),
1054 capabilities: ServerCapabilities::default(),
1055 })),
1056 server_handle: None,
1057 running: Arc::new(RwLock::new(false)),
1058 pending_request_handler: None,
1059 pending_tool_schemas: HashMap::new(),
1060 pending_capabilities: ServerCapabilities::default(),
1061 pending_task_notifications: None,
1062 #[cfg(feature = "tls")]
1063 mtls_config: None,
1064 }
1065 }
1066
1067 pub async fn set_request_handler<F>(&mut self, handler: F)
1072 where
1073 F: Fn(JsonRpcRequest) -> tokio::sync::oneshot::Receiver<JsonRpcResponse>
1074 + Send
1075 + Sync
1076 + 'static,
1077 {
1078 let mut state = self.state.write().await;
1079 state.request_handler = Some(Arc::new(move |request| {
1080 let response = handler(request);
1081 let (tx, rx) = tokio::sync::oneshot::channel();
1082 tokio::spawn(async move {
1083 let result = response.await.map_err(|error| {
1084 McpError::Internal(format!("HTTP request handler channel closed: {error}"))
1085 });
1086 let _ = tx.send(result);
1087 });
1088 rx
1089 }));
1090 }
1091
1092 #[cfg(feature = "tls")]
1094 pub fn with_mtls(mut self, config: MtlsServerConfig) -> Self {
1095 self.mtls_config = Some(config);
1096 self
1097 }
1098
1099 #[cfg(test)]
1100 pub fn get_bind_addr(&self) -> &str {
1101 &self.bind_addr
1102 }
1103
1104 #[cfg(test)]
1105 pub fn get_config(&self) -> &TransportConfig {
1106 &self.config
1107 }
1108}
1109
1110#[async_trait]
1111impl ServerTransport for HttpServerTransport {
1112 fn set_tool_schemas(&mut self, schemas: HashMap<String, Value>) -> McpResult<()> {
1113 for (name, schema) in &schemas {
1114 tool_header_mappings(schema).map_err(|error| {
1115 McpError::Validation(format!(
1116 "tool {name} has an invalid x-mcp-header schema: {error}"
1117 ))
1118 })?;
1119 }
1120 self.pending_tool_schemas = schemas;
1124 Ok(())
1125 }
1126
1127 fn set_server_capabilities(&mut self, capabilities: ServerCapabilities) -> McpResult<()> {
1128 self.pending_capabilities = capabilities;
1129 Ok(())
1130 }
1131
1132 fn set_task_notifications(
1133 &mut self,
1134 receiver: broadcast::Receiver<JsonRpcNotification>,
1135 ) -> McpResult<()> {
1136 self.pending_task_notifications = Some(receiver);
1137 Ok(())
1138 }
1139
1140 async fn start(&mut self) -> McpResult<()> {
1141 tracing::info!("Starting HTTP server on {}", self.bind_addr);
1142
1143 if let Some(handler) = self.pending_request_handler.take() {
1144 let http_handler = Arc::new(move |request: JsonRpcRequest| {
1145 let (tx, rx) = tokio::sync::oneshot::channel();
1146 let handler_future = handler(request);
1147 let parent_span = tracing::Span::current();
1148 tokio::spawn(
1149 async move {
1150 let _ = tx.send(handler_future.await);
1151 }
1152 .instrument(parent_span),
1153 );
1154 rx
1155 });
1156 self.state.write().await.request_handler = Some(http_handler);
1157 }
1158 self.state.write().await.tool_schemas = std::mem::take(&mut self.pending_tool_schemas);
1159 self.state.write().await.capabilities = std::mem::take(&mut self.pending_capabilities);
1160 if let Some(mut task_notifications) = self.pending_task_notifications.take() {
1161 let sender = self.state.read().await.notification_sender.clone();
1162 tokio::spawn(async move {
1163 loop {
1164 match task_notifications.recv().await {
1165 Ok(notification) => {
1166 let _ = sender.send(notification);
1167 }
1168 Err(broadcast::error::RecvError::Lagged(_)) => continue,
1169 Err(broadcast::error::RecvError::Closed) => break,
1170 }
1171 }
1172 });
1173 }
1174
1175 let state = self.state.clone();
1176 let bind_addr = self.bind_addr.clone();
1177 let running = self.running.clone();
1178 let _config = self.config.clone();
1179
1180 let mut app = Router::new()
1182 .route("/mcp", post(handle_mcp_request))
1183 .route("/mcp/notify", post(handle_mcp_notification))
1184 .route("/mcp/events", get(handle_sse_events))
1185 .route("/health", get(handle_health_check))
1186 .with_state(state);
1187
1188 let cors_layer = CorsLayer::new()
1190 .allow_origin(Any)
1191 .allow_methods(Any)
1192 .allow_headers(Any);
1193
1194 app = app.layer(ServiceBuilder::new().layer(cors_layer).into_inner());
1195
1196 #[cfg(feature = "tls")]
1200 let server_tls_config = self
1201 .mtls_config
1202 .as_ref()
1203 .map(MtlsServerConfig::build_rustls)
1204 .transpose()?
1205 .map(Arc::new);
1206
1207 let listener = tokio::net::TcpListener::bind(&bind_addr)
1209 .await
1210 .map_err(|e| McpError::Http(format!("Failed to bind to {bind_addr}: {e}")))?;
1211
1212 *running.write().await = true;
1213
1214 let server_handle = tokio::spawn(async move {
1215 #[cfg(feature = "tls")]
1216 if let Some(server_tls_config) = server_tls_config {
1217 let acceptor = tokio_rustls::TlsAcceptor::from(server_tls_config);
1218 loop {
1219 let (tcp_stream, peer) = match listener.accept().await {
1220 Ok(connection) => connection,
1221 Err(error) => {
1222 tracing::error!(%error, "mTLS TCP accept failed");
1223 break;
1224 }
1225 };
1226 let acceptor = acceptor.clone();
1227 let service = app.clone();
1228 tokio::spawn(async move {
1229 let tls_stream = match acceptor.accept(tcp_stream).await {
1230 Ok(stream) => stream,
1231 Err(error) => {
1232 tracing::warn!(%peer, %error, "mTLS handshake rejected");
1233 return;
1234 }
1235 };
1236 let service = TowerToHyperService::new(service);
1237 if let Err(error) = HyperServerBuilder::new(TokioExecutor::new())
1238 .serve_connection_with_upgrades(TokioIo::new(tls_stream), service)
1239 .await
1240 {
1241 tracing::debug!(%peer, %error, "mTLS HTTP connection ended");
1242 }
1243 });
1244 }
1245 return;
1246 }
1247
1248 if let Err(e) = axum::serve(listener, app).await {
1249 tracing::error!("HTTP server error: {}", e);
1250 }
1251 });
1252
1253 self.server_handle = Some(server_handle);
1254
1255 tracing::info!("HTTP server started successfully on {}", self.bind_addr);
1256 Ok(())
1257 }
1258
1259 fn set_request_handler(&mut self, handler: crate::transport::traits::ServerRequestHandler) {
1260 self.pending_request_handler = Some(handler);
1261 }
1262
1263 async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
1264 let state = self.state.read().await;
1265
1266 if state.notification_sender.send(notification).is_err() {
1267 tracing::warn!("No SSE clients connected to receive notification");
1268 }
1269
1270 Ok(())
1271 }
1272
1273 async fn stop(&mut self) -> McpResult<()> {
1274 tracing::info!("Stopping HTTP server");
1275
1276 *self.running.write().await = false;
1277
1278 if let Some(handle) = self.server_handle.take() {
1279 handle.abort();
1280 }
1281
1282 Ok(())
1283 }
1284
1285 fn is_running(&self) -> bool {
1286 self.server_handle.is_some()
1288 }
1289
1290 fn server_info(&self) -> String {
1291 format!("HTTP server transport (bind: {})", self.bind_addr)
1292 }
1293}
1294
1295async fn handle_mcp_request(
1301 State(state): State<Arc<RwLock<HttpServerState>>>,
1302 headers: HeaderMap,
1303 Json(message): Json<JsonRpcMessage>,
1304) -> Result<Response, StatusCode> {
1305 let protocol_header = headers
1306 .get(MCP_PROTOCOL_VERSION_HEADER)
1307 .and_then(|value| value.to_str().ok())
1308 .map(str::to_string);
1309 let method_header = headers
1310 .get(MCP_METHOD_HEADER)
1311 .and_then(|value| value.to_str().ok())
1312 .map(str::to_string);
1313 let name_header = headers
1314 .get(MCP_NAME_HEADER)
1315 .and_then(|value| value.to_str().ok())
1316 .map(str::to_string);
1317 let accept_header = headers
1318 .get(axum::http::header::ACCEPT)
1319 .and_then(|value| value.to_str().ok())
1320 .unwrap_or_default()
1321 .to_string();
1322 let custom_headers: HashMap<String, String> = headers
1323 .iter()
1324 .filter_map(|(name, value)| {
1325 name.as_str()
1326 .to_ascii_lowercase()
1327 .starts_with("mcp-param-")
1328 .then(|| {
1329 value
1330 .to_str()
1331 .ok()
1332 .map(|value| (name.as_str().to_string(), value.to_string()))
1333 })
1334 .flatten()
1335 })
1336 .collect();
1337 let dispatch = async move {
1338 match message {
1339 JsonRpcMessage::Request(request) => {
1340 let is_modern = request_protocol_version(&request).is_some();
1341 if let Err(error) = validate_http_headers(
1342 &request,
1343 protocol_header.as_deref(),
1344 method_header.as_deref(),
1345 name_header.as_deref(),
1346 ) {
1347 let (code, data) = json_rpc_error_details(&error);
1348 let body = JsonRpcMessage::Error(JsonRpcError::error(
1349 request.id,
1350 code,
1351 error.to_string(),
1352 data,
1353 ));
1354 return Ok((StatusCode::BAD_REQUEST, Json(body)).into_response());
1355 }
1356 if is_modern && request.method == methods::SUBSCRIPTIONS_LISTEN {
1357 if !accept_header
1358 .split(',')
1359 .any(|value| value.trim().starts_with("text/event-stream"))
1360 {
1361 let error = McpError::HeaderMismatch(
1362 "subscriptions/listen requires Accept: text/event-stream".to_string(),
1363 );
1364 let (code, data) = json_rpc_error_details(&error);
1365 let body = JsonRpcMessage::Error(JsonRpcError::error(
1366 request.id,
1367 code,
1368 error.to_string(),
1369 data,
1370 ));
1371 return Ok((StatusCode::BAD_REQUEST, Json(body)).into_response());
1372 }
1373 return handle_subscription_stream(state, request).await;
1374 }
1375 if is_modern && request.method == methods::TOOLS_CALL {
1376 let params = request.params.as_ref().and_then(Value::as_object);
1377 let tool_name = params
1378 .and_then(|params| params.get("name"))
1379 .and_then(Value::as_str);
1380 let arguments = params
1381 .and_then(|params| params.get("arguments"))
1382 .unwrap_or(&Value::Null);
1383 let schema = if let Some(tool_name) = tool_name {
1384 state.read().await.tool_schemas.get(tool_name).cloned()
1385 } else {
1386 None
1387 };
1388 let validation = match schema {
1389 Some(schema) => {
1390 validate_tool_call_headers(&schema, arguments, &custom_headers)
1391 }
1392 None => Ok(()),
1393 };
1394 if let Err(error) = validation {
1395 let (code, data) = json_rpc_error_details(&error);
1396 let body = JsonRpcMessage::Error(JsonRpcError::error(
1397 request.id,
1398 code,
1399 error.to_string(),
1400 data,
1401 ));
1402 return Ok((StatusCode::BAD_REQUEST, Json(body)).into_response());
1403 }
1404 }
1405 handle_mcp_jsonrpc_request(state, request)
1406 .await
1407 .map(|message| {
1408 let status = match &message {
1409 JsonRpcMessage::Error(error)
1410 if matches!(
1411 error.error.code,
1412 HEADER_MISMATCH
1413 | UNSUPPORTED_PROTOCOL_VERSION
1414 | crate::protocol::MISSING_REQUIRED_CLIENT_CAPABILITY
1415 | error_codes::INVALID_REQUEST
1416 | error_codes::INVALID_PARAMS
1417 ) =>
1418 {
1419 StatusCode::BAD_REQUEST
1420 }
1421 JsonRpcMessage::Error(error)
1422 if is_modern
1423 && error.error.code == error_codes::METHOD_NOT_FOUND =>
1424 {
1425 StatusCode::NOT_FOUND
1426 }
1427 JsonRpcMessage::Error(error) if error.error.code == FORBIDDEN_ERROR => {
1428 StatusCode::FORBIDDEN
1429 }
1430 JsonRpcMessage::Error(error)
1431 if error.error.code == RATE_LIMITED_ERROR =>
1432 {
1433 StatusCode::TOO_MANY_REQUESTS
1434 }
1435 _ => StatusCode::OK,
1436 };
1437 (status, Json(message)).into_response()
1438 })
1439 }
1440 JsonRpcMessage::Notification(notification) => {
1441 handle_mcp_jsonrpc_notification(state, notification).await?;
1442 Ok(StatusCode::ACCEPTED.into_response())
1443 }
1444 JsonRpcMessage::Response(_) | JsonRpcMessage::Error(_) => Err(StatusCode::BAD_REQUEST),
1445 }
1446 };
1447
1448 #[cfg(feature = "otel")]
1449 {
1450 use tracing_opentelemetry::OpenTelemetrySpanExt;
1451
1452 let parent = opentelemetry::global::get_text_map_propagator(|propagator| {
1453 propagator.extract(&HeaderExtractor(&headers))
1454 });
1455 let span = tracing::info_span!("mcp.http", otel.kind = "server");
1456 let _ = span.set_parent(parent);
1457 dispatch.instrument(span).await
1458 }
1459
1460 #[cfg(not(feature = "otel"))]
1461 {
1462 dispatch.await
1463 }
1464}
1465
1466async fn handle_subscription_stream(
1467 state: Arc<RwLock<HttpServerState>>,
1468 request: JsonRpcRequest,
1469) -> Result<Response, StatusCode> {
1470 let context = modern_request_context(&request).map_err(|_| StatusCode::BAD_REQUEST)?;
1471 let context = context.ok_or(StatusCode::BAD_REQUEST)?;
1472 let params: SubscriptionsListenParams =
1473 request
1474 .params
1475 .clone()
1476 .ok_or(StatusCode::BAD_REQUEST)
1477 .and_then(|value| serde_json::from_value(value).map_err(|_| StatusCode::BAD_REQUEST))?;
1478 if params.notifications.requests_tasks() && !has_tasks_extension(&context.client_capabilities) {
1479 let error = McpError::MissingRequiredClientCapability(serde_json::json!({
1480 "extensions": {(TASKS_EXTENSION_ID): {}}
1481 }));
1482 let (code, data) = json_rpc_error_details(&error);
1483 let body = JsonRpcMessage::Error(JsonRpcError::error(
1484 request.id,
1485 code,
1486 error.to_string(),
1487 data,
1488 ));
1489 return Ok((StatusCode::BAD_REQUEST, Json(body)).into_response());
1490 }
1491
1492 let state_guard = state.read().await;
1493 let capabilities = &state_guard.capabilities;
1494 let tasks_enabled = capabilities
1495 .extensions
1496 .as_ref()
1497 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID));
1498 let accepted = SubscriptionFilter {
1499 tools_list_changed: (params.notifications.tools_list_changed == Some(true)
1500 && capabilities
1501 .tools
1502 .as_ref()
1503 .is_some_and(|value| value.list_changed == Some(true)))
1504 .then_some(true),
1505 prompts_list_changed: (params.notifications.prompts_list_changed == Some(true)
1506 && capabilities
1507 .prompts
1508 .as_ref()
1509 .is_some_and(|value| value.list_changed == Some(true)))
1510 .then_some(true),
1511 resources_list_changed: (params.notifications.resources_list_changed == Some(true)
1512 && capabilities
1513 .resources
1514 .as_ref()
1515 .is_some_and(|value| value.list_changed == Some(true)))
1516 .then_some(true),
1517 resource_subscriptions: if capabilities
1518 .resources
1519 .as_ref()
1520 .is_some_and(|value| value.subscribe == Some(true))
1521 {
1522 params.notifications.resource_subscriptions.clone()
1523 } else {
1524 Vec::new()
1525 },
1526 task_ids: if tasks_enabled {
1527 params.notifications.task_ids.clone()
1528 } else {
1529 Vec::new()
1530 },
1531 };
1532 let receiver = state_guard.notification_sender.subscribe();
1533 drop(state_guard);
1534
1535 let subscription_id = request.id.clone();
1536 let mut ack_meta = HashMap::new();
1537 ack_meta.insert(
1538 SUBSCRIPTION_ID_META_KEY.to_string(),
1539 subscription_id.clone(),
1540 );
1541 let acknowledgement = JsonRpcNotification::new(
1542 methods::SUBSCRIPTIONS_ACKNOWLEDGED.to_string(),
1543 Some(SubscriptionsAcknowledgedParams {
1544 notifications: accepted.clone(),
1545 meta: ack_meta,
1546 }),
1547 )
1548 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
1549
1550 let stream = futures::stream::unfold(
1551 (Some(acknowledgement), receiver, accepted, subscription_id),
1552 |(first, mut receiver, filter, subscription_id)| async move {
1553 if let Some(notification) = first {
1554 let data =
1555 serde_json::to_string(¬ification).unwrap_or_else(|_| "{}".to_string());
1556 return Some((
1557 Ok::<Event, Infallible>(Event::default().data(data)),
1558 (None, receiver, filter, subscription_id),
1559 ));
1560 }
1561 loop {
1562 let mut notification = match receiver.recv().await {
1563 Ok(notification) => notification,
1564 Err(broadcast::error::RecvError::Lagged(_)) => continue,
1565 Err(broadcast::error::RecvError::Closed) => return None,
1566 };
1567 if !filter.matches(¬ification.method, notification.params.as_ref()) {
1568 continue;
1569 }
1570 let params = notification
1571 .params
1572 .get_or_insert_with(|| Value::Object(serde_json::Map::new()));
1573 let Some(object) = params.as_object_mut() else {
1574 continue;
1575 };
1576 let meta = object
1577 .entry("_meta")
1578 .or_insert_with(|| Value::Object(serde_json::Map::new()));
1579 let Some(meta) = meta.as_object_mut() else {
1580 continue;
1581 };
1582 meta.insert(
1583 SUBSCRIPTION_ID_META_KEY.to_string(),
1584 subscription_id.clone(),
1585 );
1586 let data =
1587 serde_json::to_string(¬ification).unwrap_or_else(|_| "{}".to_string());
1588 return Some((
1589 Ok::<Event, Infallible>(Event::default().data(data)),
1590 (None, receiver, filter, subscription_id),
1591 ));
1592 }
1593 },
1594 );
1595
1596 Ok(Sse::new(stream)
1597 .keep_alive(
1598 axum::response::sse::KeepAlive::new()
1599 .interval(Duration::from_secs(30))
1600 .text("keep-alive"),
1601 )
1602 .into_response())
1603}
1604
1605async fn handle_mcp_jsonrpc_request(
1606 state: Arc<RwLock<HttpServerState>>,
1607 request: JsonRpcRequest,
1608) -> Result<JsonRpcMessage, StatusCode> {
1609 let state_guard = state.read().await;
1610
1611 if let Some(ref handler) = state_guard.request_handler {
1612 let request_id = request.id.clone();
1613 let response_rx = handler(request);
1614 drop(state_guard); match response_rx.await {
1617 Ok(Ok(response)) => Ok(JsonRpcMessage::Response(response)),
1618 Ok(Err(error)) => {
1619 let (code, data) = match &error {
1620 McpError::Forbidden(_) => (FORBIDDEN_ERROR, None),
1621 McpError::RateLimited { retry_after_ms } => (
1622 RATE_LIMITED_ERROR,
1623 Some(serde_json::json!({"retryAfterMs": retry_after_ms})),
1624 ),
1625 _ => json_rpc_error_details(&error),
1626 };
1627 Ok(JsonRpcMessage::Error(JsonRpcError::error(
1628 request_id,
1629 code,
1630 error.to_string(),
1631 data,
1632 )))
1633 }
1634 Err(_) => Err(StatusCode::INTERNAL_SERVER_ERROR),
1635 }
1636 } else {
1637 let error_response = JsonRpcError::error(
1638 request.id,
1639 error_codes::METHOD_NOT_FOUND,
1640 "No request handler configured".to_string(),
1641 None,
1642 );
1643 Ok(JsonRpcMessage::Error(error_response))
1644 }
1645}
1646
1647async fn handle_mcp_jsonrpc_notification(
1648 state: Arc<RwLock<HttpServerState>>,
1649 notification: JsonRpcNotification,
1650) -> Result<(), StatusCode> {
1651 if !is_supported_http_notification(¬ification) {
1652 return Err(StatusCode::BAD_REQUEST);
1653 }
1654
1655 let state_guard = state.read().await;
1656 if state_guard.notification_sender.send(notification).is_err() {
1657 tracing::debug!("No SSE clients connected to receive notification");
1658 }
1659
1660 Ok(())
1661}
1662
1663fn is_supported_http_notification(notification: &JsonRpcNotification) -> bool {
1664 if notification.jsonrpc != "2.0" {
1665 return false;
1666 }
1667
1668 matches!(
1669 notification.method.as_str(),
1670 methods::INITIALIZED
1671 | methods::TOOLS_LIST_CHANGED
1672 | methods::RESOURCES_UPDATED
1673 | methods::RESOURCES_LIST_CHANGED
1674 | methods::PROMPTS_LIST_CHANGED
1675 | methods::ROOTS_LIST_CHANGED
1676 | methods::ELICITATION_COMPLETE
1677 | methods::TASKS_STATUS
1678 | methods::TASKS_STATUS_UPDATE
1679 | methods::LOGGING_MESSAGE
1680 | methods::PROGRESS
1681 | methods::CANCELLED
1682 )
1683}
1684
1685async fn handle_mcp_notification(Json(_notification): Json<JsonRpcNotification>) -> StatusCode {
1687 StatusCode::OK
1689}
1690
1691#[cfg(feature = "sse")]
1693async fn handle_sse_events(
1694 State(state): State<Arc<RwLock<HttpServerState>>>,
1695) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {
1696 let state_guard = state.read().await;
1697 let receiver = state_guard.notification_sender.subscribe();
1698 drop(state_guard);
1699
1700 let stream = BroadcastStream::new(receiver).map(|result| {
1701 match result {
1702 Ok(notification) => match serde_json::to_string(¬ification) {
1703 Ok(json) => Ok(Event::default().data(json)),
1704 Err(e) => {
1705 tracing::error!("Failed to serialize notification: {}", e);
1706 Ok(Event::default().data("{}"))
1707 }
1708 },
1709 Err(_) => Ok(Event::default().data("{}")), }
1711 });
1712
1713 Sse::new(stream).keep_alive(
1714 axum::response::sse::KeepAlive::new()
1715 .interval(Duration::from_secs(30))
1716 .text("keep-alive"),
1717 )
1718}
1719
1720#[cfg(not(feature = "sse"))]
1722async fn handle_sse_events(_state: State<Arc<RwLock<HttpServerState>>>) -> StatusCode {
1723 StatusCode::NOT_IMPLEMENTED
1724}
1725
1726async fn handle_health_check() -> Json<Value> {
1728 let timestamp = chrono::Utc::now().to_rfc3339();
1729
1730 Json(serde_json::json!({
1731 "status": "healthy",
1732 "transport": "http",
1733 "timestamp": timestamp
1734 }))
1735}
1736
1737#[cfg(test)]
1738mod tests {
1739 use super::*;
1740 use crate::protocol::methods;
1741 use wiremock::matchers::{header, method, path};
1742 use wiremock::{Mock, MockServer, ResponseTemplate};
1743
1744 #[test]
1745 fn parses_json_rpc_result_from_standard_post_sse() {
1746 let body = b"event: message\r\ndata: {\"jsonrpc\":\"2.0\",\"id\":7,\"result\":{\"ok\":true}}\r\n\r\n";
1747 let parsed = parse_sse_response(body, &serde_json::json!(7)).unwrap();
1748 assert_eq!(parsed["result"]["ok"], true);
1749 }
1750
1751 #[tokio::test]
1752 async fn test_http_client_creation() {
1753 let transport = HttpClientTransport::new("http://localhost:3000", None).await;
1754 assert!(transport.is_ok());
1755
1756 let transport = transport.unwrap();
1757 assert!(transport.is_connected());
1758 assert_eq!(transport.base_url, "http://localhost:3000");
1759 }
1760
1761 #[tokio::test]
1762 async fn test_http_server_creation() {
1763 let transport = HttpServerTransport::new("127.0.0.1:0");
1764 assert_eq!(transport.bind_addr, "127.0.0.1:0");
1765 assert!(!transport.is_running());
1766 }
1767
1768 #[test]
1769 fn test_http_server_with_config() {
1770 let config = TransportConfig {
1771 compression: true,
1772 ..Default::default()
1773 };
1774
1775 let transport = HttpServerTransport::with_config("0.0.0.0:8080", config);
1776 assert_eq!(transport.bind_addr, "0.0.0.0:8080");
1777 assert!(transport.config.compression);
1778 }
1779
1780 #[tokio::test]
1781 async fn test_http_client_with_sse() {
1782 let transport = HttpClientTransport::new(
1783 "http://localhost:3000",
1784 Some("http://localhost:3000/events"),
1785 )
1786 .await;
1787
1788 assert!(transport.is_ok());
1789 let transport = transport.unwrap();
1790 assert!(transport.sse_url.is_some());
1791 assert_eq!(transport.sse_url.unwrap(), "http://localhost:3000/events");
1792 }
1793
1794 #[tokio::test]
1796 async fn test_request_id_generation_sequence() {
1797 let transport = HttpClientTransport::new("http://localhost:3000", None)
1798 .await
1799 .unwrap();
1800
1801 let id1 = transport.next_request_id().await;
1802 let id2 = transport.next_request_id().await;
1803 let id3 = transport.next_request_id().await;
1804
1805 assert_eq!(id1, 1);
1806 assert_eq!(id2, 2);
1807 assert_eq!(id3, 3);
1808 }
1809
1810 #[tokio::test]
1811 async fn test_request_tracking_complete() {
1812 let transport = HttpClientTransport::new("http://localhost:3000", None)
1813 .await
1814 .unwrap();
1815
1816 assert_eq!(transport.active_request_count().await, 0);
1818
1819 let request_ids = vec![
1821 Value::from(123),
1822 Value::String("string-id".to_string()),
1823 Value::Null,
1824 Value::Array(vec![Value::from(1), Value::from(2)]),
1825 ];
1826
1827 for id in &request_ids {
1828 transport.track_request(id).await;
1829 }
1830 assert_eq!(transport.active_request_count().await, request_ids.len());
1831
1832 for id in &request_ids {
1834 transport.untrack_request(id).await;
1835 }
1836 assert_eq!(transport.active_request_count().await, 0);
1837
1838 transport.untrack_request(&Value::from(999)).await;
1840 assert_eq!(transport.active_request_count().await, 0);
1841 }
1842
1843 #[tokio::test]
1844 async fn test_connection_state_management() {
1845 let mut transport = HttpClientTransport::new("http://localhost:3000", None)
1846 .await
1847 .unwrap();
1848
1849 assert!(transport.is_connected());
1851 assert!(transport.has_notification_receiver());
1852
1853 let info_before = transport.connection_info();
1854 assert!(info_before.contains("Connected"));
1855
1856 let result = transport.close().await;
1858 assert!(result.is_ok());
1859
1860 assert!(!transport.is_connected());
1862 assert!(!transport.has_notification_receiver());
1863
1864 let info_after = transport.connection_info();
1865 assert!(info_after.contains("Disconnected"));
1866 }
1867
1868 #[tokio::test]
1869 async fn test_receive_notification_states() {
1870 let mut transport = HttpClientTransport::new("http://localhost:3000", None)
1871 .await
1872 .unwrap();
1873
1874 let result = transport.receive_notification().await;
1877 assert!(result.is_err());
1878 assert!(result.unwrap_err().to_string().contains("disconnected"));
1879
1880 transport.close().await.unwrap();
1882 let result = transport.receive_notification().await;
1883 assert!(result.is_ok());
1884 assert!(result.unwrap().is_none());
1885
1886 let result2 = transport.receive_notification().await;
1888 assert!(result2.is_ok());
1889 assert!(result2.unwrap().is_none());
1890 }
1891
1892 #[tokio::test]
1893 async fn test_http_server_lifecycle_complete() {
1894 let mut transport = HttpServerTransport::new("127.0.0.1:0");
1895
1896 assert_eq!(transport.get_bind_addr(), "127.0.0.1:0");
1898 assert!(!transport.is_running());
1899
1900 let info = transport.server_info();
1901 assert!(info.contains("HTTP server transport"));
1902 assert!(info.contains("127.0.0.1:0"));
1903
1904 let result = transport.start().await;
1906 assert!(result.is_ok());
1907 assert!(transport.is_running());
1908
1909 let notification = JsonRpcNotification {
1911 jsonrpc: "2.0".to_string(),
1912 method: "test_notification".to_string(),
1913 params: Some(serde_json::json!({"test": true})),
1914 };
1915 let result = transport.send_notification(notification).await;
1916 assert!(result.is_ok());
1917
1918 let result = transport.stop().await;
1920 assert!(result.is_ok());
1921 assert!(!transport.is_running());
1922
1923 let result = transport.stop().await;
1925 assert!(result.is_ok());
1926 }
1927
1928 #[tokio::test]
1929 async fn test_http_server_request_handler() {
1930 let mut transport = HttpServerTransport::new("127.0.0.1:0");
1931
1932 let handler = |request: JsonRpcRequest| {
1933 let (tx, rx) = tokio::sync::oneshot::channel();
1934 let response = JsonRpcResponse {
1935 jsonrpc: "2.0".to_string(),
1936 id: request.id,
1937 result: Some(serde_json::json!({
1938 "method_received": request.method,
1939 "handled": true
1940 })),
1941 };
1942 let _ = tx.send(response);
1943 rx
1944 };
1945
1946 transport.set_request_handler(handler).await;
1947 }
1949
1950 #[tokio::test]
1951 async fn test_http_server_with_custom_config() {
1952 let mut config = TransportConfig {
1953 compression: true,
1954 ..Default::default()
1955 };
1956 config
1957 .headers
1958 .insert("Server".to_string(), "MCP-Test/1.0".to_string());
1959
1960 let transport = HttpServerTransport::with_config("0.0.0.0:8080", config);
1961
1962 assert_eq!(transport.get_bind_addr(), "0.0.0.0:8080");
1963 assert!(transport.get_config().compression);
1964 assert_eq!(
1965 transport.get_config().headers.get("Server"),
1966 Some(&"MCP-Test/1.0".to_string())
1967 );
1968 }
1969
1970 #[tokio::test]
1971 async fn test_http_client_with_custom_config() {
1972 let mut config = TransportConfig {
1973 read_timeout_ms: Some(5000),
1974 connect_timeout_ms: Some(2000),
1975 write_timeout_ms: Some(3000),
1976 ..Default::default()
1977 };
1978 config
1979 .headers
1980 .insert("X-Custom-Header".to_string(), "test-value".to_string());
1981 config
1982 .headers
1983 .insert("Authorization".to_string(), "Bearer token123".to_string());
1984
1985 let transport = HttpClientTransport::with_config(
1986 "http://localhost:3000",
1987 Some("http://localhost:3000/events"),
1988 config,
1989 )
1990 .await;
1991
1992 assert!(transport.is_ok());
1993 let transport = transport.unwrap();
1994 assert_eq!(transport.config.read_timeout_ms, Some(5000));
1995 assert_eq!(transport.config.connect_timeout_ms, Some(2000));
1996 assert_eq!(transport.config.write_timeout_ms, Some(3000));
1997 assert!(transport.sse_url.is_some());
1998 }
1999
2000 #[tokio::test]
2002 async fn test_handle_health_check() {
2003 let result = handle_health_check().await;
2004
2005 let Json(health_data) = result;
2006 assert_eq!(health_data["status"], "healthy");
2007 assert_eq!(health_data["transport"], "http");
2008 assert!(health_data["timestamp"].is_string());
2009 }
2010
2011 #[tokio::test]
2012 async fn test_handle_mcp_notification() {
2013 let notification = JsonRpcNotification {
2014 jsonrpc: "2.0".to_string(),
2015 method: "test_notification".to_string(),
2016 params: Some(serde_json::json!({"test": "notification"})),
2017 };
2018 let json_notification = Json(notification);
2019
2020 let result = handle_mcp_notification(json_notification).await;
2021
2022 assert_eq!(result, StatusCode::OK);
2024 }
2025
2026 #[tokio::test]
2027 async fn test_handle_mcp_request_accepts_initialized_notification() {
2028 let (notification_sender, mut notification_receiver) = broadcast::channel(100);
2029 let state = Arc::new(RwLock::new(HttpServerState {
2030 notification_sender,
2031 request_handler: None,
2032 tool_schemas: HashMap::new(),
2033 capabilities: ServerCapabilities::default(),
2034 }));
2035 let message = serde_json::from_value::<JsonRpcMessage>(serde_json::json!({
2036 "jsonrpc": "2.0",
2037 "method": methods::INITIALIZED,
2038 "params": {}
2039 }))
2040 .unwrap();
2041
2042 let response = handle_mcp_request(State(state), HeaderMap::new(), Json(message))
2043 .await
2044 .unwrap();
2045
2046 assert_eq!(response.status(), StatusCode::ACCEPTED);
2047 let received = notification_receiver.recv().await.unwrap();
2048 assert_eq!(received.method, methods::INITIALIZED);
2049 }
2050
2051 #[tokio::test]
2052 async fn test_handle_mcp_request_rejects_unknown_notification() {
2053 let (notification_sender, _) = broadcast::channel(100);
2054 let state = Arc::new(RwLock::new(HttpServerState {
2055 notification_sender,
2056 request_handler: None,
2057 tool_schemas: HashMap::new(),
2058 capabilities: ServerCapabilities::default(),
2059 }));
2060 let message = serde_json::from_value::<JsonRpcMessage>(serde_json::json!({
2061 "jsonrpc": "2.0",
2062 "method": "notifications/unknown",
2063 "params": {}
2064 }))
2065 .unwrap();
2066
2067 let result = handle_mcp_request(State(state), HeaderMap::new(), Json(message)).await;
2068
2069 assert!(matches!(result, Err(StatusCode::BAD_REQUEST)));
2070 }
2071
2072 #[cfg(not(feature = "sse"))]
2073 #[tokio::test]
2074 async fn test_handle_sse_events_not_implemented() {
2075 let (notification_sender, _) = broadcast::channel(100);
2076
2077 let state = Arc::new(RwLock::new(HttpServerState {
2078 notification_sender,
2079 request_handler: None,
2080 tool_schemas: HashMap::new(),
2081 capabilities: ServerCapabilities::default(),
2082 }));
2083
2084 let state_extract = State(state);
2085
2086 let result = handle_sse_events(state_extract).await;
2087
2088 assert_eq!(result, StatusCode::NOT_IMPLEMENTED);
2090 }
2091
2092 #[tokio::test]
2094 async fn test_transport_config_variations() {
2095 let default_config = TransportConfig::default();
2097 assert_eq!(default_config.read_timeout_ms, Some(60_000));
2098 assert_eq!(default_config.write_timeout_ms, Some(30_000));
2099 assert_eq!(default_config.connect_timeout_ms, Some(30_000));
2100 assert!(default_config.headers.is_empty());
2101
2102 let mut full_config = TransportConfig {
2104 read_timeout_ms: Some(10000),
2105 write_timeout_ms: Some(5000),
2106 connect_timeout_ms: Some(3000),
2107 compression: true,
2108 ..Default::default()
2109 };
2110 full_config
2111 .headers
2112 .insert("Test-Header".to_string(), "test-value".to_string());
2113
2114 let transport =
2115 HttpClientTransport::with_config("http://localhost:3000", None, full_config)
2116 .await
2117 .unwrap();
2118
2119 assert_eq!(transport.config.read_timeout_ms, Some(10000));
2120 assert_eq!(transport.config.write_timeout_ms, Some(5000));
2121 assert_eq!(transport.config.connect_timeout_ms, Some(3000));
2122 assert!(transport.config.compression);
2123 }
2124
2125 #[tokio::test]
2126 async fn test_sse_url_variations() {
2127 let transport1 = HttpClientTransport::new(
2129 "http://localhost:3000",
2130 Some("http://localhost:3000/events"),
2131 )
2132 .await
2133 .unwrap();
2134 assert!(transport1.sse_url.is_some());
2135 assert_eq!(
2136 transport1.sse_url.as_ref().unwrap(),
2137 "http://localhost:3000/events"
2138 );
2139
2140 let transport2 = HttpClientTransport::new(
2142 "http://localhost:3000",
2143 Some("http://localhost:3000/events"),
2144 )
2145 .await
2146 .unwrap();
2147 assert!(transport2.sse_url.is_some());
2148
2149 let transport3 = HttpClientTransport::new("http://localhost:3000", None::<&str>)
2151 .await
2152 .unwrap();
2153 assert!(transport3.sse_url.is_none());
2154
2155 let info1 = transport1.connection_info();
2157 assert!(info1.contains("http://localhost:3000/events"));
2158
2159 let info3 = transport3.connection_info();
2160 assert!(info3.contains("sse: None"));
2161 }
2162
2163 #[tokio::test]
2164 async fn test_concurrent_request_id_generation() {
2165 let transport = std::sync::Arc::new(
2166 HttpClientTransport::new("http://localhost:3000", None)
2167 .await
2168 .unwrap(),
2169 );
2170
2171 let mut handles = vec![];
2172
2173 for _ in 0..3 {
2175 let transport_clone = transport.clone();
2176 let handle = tokio::spawn(async move {
2177 let mut ids = vec![];
2178 for _ in 0..3 {
2179 ids.push(transport_clone.next_request_id().await);
2180 }
2181 ids
2182 });
2183 handles.push(handle);
2184 }
2185
2186 let mut all_ids = vec![];
2187 for handle in handles {
2188 let ids = handle.await.unwrap();
2189 all_ids.extend(ids);
2190 }
2191
2192 all_ids.sort();
2194 let mut unique_ids = all_ids.clone();
2195 unique_ids.dedup();
2196
2197 assert_eq!(all_ids.len(), unique_ids.len());
2198 assert_eq!(all_ids.len(), 9); }
2200
2201 #[tokio::test]
2202 async fn test_server_bind_addresses() {
2203 let test_cases = vec!["127.0.0.1:0", "0.0.0.0:8080", "localhost:9000"];
2204
2205 for addr in test_cases {
2206 let server = HttpServerTransport::new(addr);
2207 assert_eq!(server.get_bind_addr(), addr);
2208 assert!(!server.is_running());
2209
2210 let info = server.server_info();
2211 assert!(info.contains("HTTP server transport"));
2212 assert!(info.contains(addr));
2213 }
2214 }
2215
2216 #[tokio::test]
2218 async fn test_transport_send_request_with_mock() {
2219 let mock_server = MockServer::start().await;
2220
2221 let expected_response = JsonRpcResponse {
2223 jsonrpc: "2.0".to_string(),
2224 id: Value::from(42),
2225 result: Some(serde_json::json!({
2226 "capabilities": {
2227 "tools": true,
2228 "resources": true
2229 }
2230 })),
2231 };
2232
2233 Mock::given(method("POST"))
2234 .and(path("/mcp"))
2235 .and(header("content-type", "application/json"))
2236 .respond_with(ResponseTemplate::new(200).set_body_json(&expected_response))
2237 .mount(&mock_server)
2238 .await;
2239
2240 let mut transport = HttpClientTransport::new(mock_server.uri(), None)
2241 .await
2242 .unwrap();
2243
2244 let request = JsonRpcRequest {
2245 jsonrpc: "2.0".to_string(),
2246 id: Value::from(42),
2247 method: "initialize".to_string(),
2248 params: Some(serde_json::json!({
2249 "protocolVersion": "2024-11-05",
2250 "capabilities": {}
2251 })),
2252 };
2253
2254 let result = transport.send_request(request).await;
2255
2256 assert!(result.is_ok());
2257 let response = result.unwrap();
2258 assert_eq!(response.id, Value::from(42));
2259 assert_eq!(response.jsonrpc, "2.0");
2260 assert!(response.result.is_some());
2261 }
2262
2263 #[tokio::test]
2264 async fn test_transport_send_notification_with_mock() {
2265 let mock_server = MockServer::start().await;
2266
2267 Mock::given(method("POST"))
2268 .and(path("/mcp"))
2269 .and(header("content-type", "application/json"))
2270 .respond_with(ResponseTemplate::new(200))
2271 .mount(&mock_server)
2272 .await;
2273
2274 let mut transport = HttpClientTransport::new(mock_server.uri(), None)
2275 .await
2276 .unwrap();
2277
2278 let notification = JsonRpcNotification {
2279 jsonrpc: "2.0".to_string(),
2280 method: "initialized".to_string(),
2281 params: Some(serde_json::json!({})),
2282 };
2283
2284 let result = transport.send_notification(notification).await;
2285 assert!(result.is_ok());
2286 }
2287
2288 #[tokio::test]
2289 async fn test_transport_request_auto_id() {
2290 let mock_server = MockServer::start().await;
2291
2292 Mock::given(method("POST"))
2293 .and(path("/mcp"))
2294 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
2295 "jsonrpc": "2.0",
2296 "id": 1,
2297 "result": {"status": "ok"}
2298 })))
2299 .mount(&mock_server)
2300 .await;
2301
2302 let mut transport = HttpClientTransport::new(mock_server.uri(), None)
2303 .await
2304 .unwrap();
2305
2306 let request = JsonRpcRequest {
2308 jsonrpc: "2.0".to_string(),
2309 id: Value::Null,
2310 method: "ping".to_string(),
2311 params: None,
2312 };
2313
2314 let result = transport.send_request(request).await;
2315 assert!(result.is_ok());
2316 let response = result.unwrap();
2317 assert_eq!(response.id, Value::from(1));
2318 }
2319
2320 #[tokio::test]
2321 async fn test_transport_error_scenarios() {
2322 let mock_server = MockServer::start().await;
2323
2324 Mock::given(method("POST"))
2326 .and(path("/mcp"))
2327 .respond_with(ResponseTemplate::new(500).set_body_string("Internal Server Error"))
2328 .mount(&mock_server)
2329 .await;
2330
2331 let mut transport = HttpClientTransport::new(mock_server.uri(), None)
2332 .await
2333 .unwrap();
2334
2335 let request = JsonRpcRequest {
2336 jsonrpc: "2.0".to_string(),
2337 id: Value::from(1),
2338 method: "test".to_string(),
2339 params: None,
2340 };
2341
2342 let result = transport.send_request(request).await;
2343 assert!(result.is_err());
2344
2345 if let Err(McpError::Http(msg)) = result {
2346 assert!(msg.contains("HTTP error: 500"));
2347 } else {
2348 panic!("Expected HTTP error");
2349 }
2350 }
2351
2352 #[tokio::test]
2353 async fn test_transport_notification_error() {
2354 let mock_server = MockServer::start().await;
2355
2356 Mock::given(method("POST"))
2357 .and(path("/mcp"))
2358 .respond_with(ResponseTemplate::new(400).set_body_string("Bad Request"))
2359 .mount(&mock_server)
2360 .await;
2361
2362 let mut transport = HttpClientTransport::new(mock_server.uri(), None)
2363 .await
2364 .unwrap();
2365
2366 let notification = JsonRpcNotification {
2367 jsonrpc: "2.0".to_string(),
2368 method: "test_notification".to_string(),
2369 params: None,
2370 };
2371
2372 let result = transport.send_notification(notification).await;
2373 assert!(result.is_err());
2374
2375 if let Err(McpError::Http(msg)) = result {
2376 assert!(msg.contains("HTTP notification error: 400"));
2377 } else {
2378 panic!("Expected HTTP notification error");
2379 }
2380 }
2381
2382 #[tokio::test]
2383 async fn test_transport_connection_failure() {
2384 let mut transport = HttpClientTransport::new("http://127.0.0.1:1", None)
2386 .await
2387 .unwrap();
2388
2389 let request = JsonRpcRequest {
2390 jsonrpc: "2.0".to_string(),
2391 id: Value::from(1),
2392 method: "test".to_string(),
2393 params: None,
2394 };
2395
2396 let result = transport.send_request(request).await;
2397 assert!(result.is_err());
2398 assert!(result.is_err());
2400 }
2401
2402 #[tokio::test]
2403 async fn test_transport_invalid_json_response() {
2404 let mock_server = MockServer::start().await;
2405
2406 Mock::given(method("POST"))
2407 .and(path("/mcp"))
2408 .respond_with(ResponseTemplate::new(200).set_body_string("not valid json"))
2409 .mount(&mock_server)
2410 .await;
2411
2412 let mut transport = HttpClientTransport::new(mock_server.uri(), None)
2413 .await
2414 .unwrap();
2415
2416 let request = JsonRpcRequest {
2417 jsonrpc: "2.0".to_string(),
2418 id: Value::from(1),
2419 method: "test".to_string(),
2420 params: None,
2421 };
2422
2423 let result = transport.send_request(request).await;
2424 assert!(result.is_err());
2425
2426 if let Err(McpError::Connection(msg)) = result {
2427 assert!(msg.contains("Request serialization failed"));
2428 } else {
2429 assert!(result.is_err());
2431 }
2432 }
2433
2434 #[tokio::test]
2435 async fn test_transport_response_id_mismatch() {
2436 let mock_server = MockServer::start().await;
2437
2438 Mock::given(method("POST"))
2439 .and(path("/mcp"))
2440 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
2441 "jsonrpc": "2.0",
2442 "id": 999, "result": {"success": true}
2444 })))
2445 .mount(&mock_server)
2446 .await;
2447
2448 let mut transport = HttpClientTransport::new(mock_server.uri(), None)
2449 .await
2450 .unwrap();
2451
2452 let request = JsonRpcRequest {
2453 jsonrpc: "2.0".to_string(),
2454 id: Value::from(1),
2455 method: "test".to_string(),
2456 params: None,
2457 };
2458
2459 let result = transport.send_request(request).await;
2460 assert!(result.is_err());
2461
2462 if let Err(McpError::Http(msg)) = result {
2463 assert!(msg.contains("Response ID") && msg.contains("does not match request ID"));
2464 } else {
2465 panic!("Expected HTTP error for ID mismatch");
2466 }
2467 }
2468
2469 #[tokio::test]
2470 async fn http_server_installs_and_runs_the_mcp_request_handler() {
2471 use crate::server::McpServer;
2472
2473 let reserved = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
2474 let address = reserved.local_addr().unwrap();
2475 drop(reserved);
2476
2477 let mut server = McpServer::create("http-e2e", "1.0.0");
2478 server
2479 .start(HttpServerTransport::new(address.to_string()))
2480 .await
2481 .unwrap();
2482
2483 let mut client = HttpClientTransport::new(format!("http://{address}"), None)
2484 .await
2485 .unwrap();
2486 let response = client
2487 .send_request(
2488 JsonRpcRequest::new(Value::from(42), methods::PING.to_string(), None::<Value>)
2489 .unwrap(),
2490 )
2491 .await
2492 .unwrap();
2493 assert_eq!(response.id, Value::from(42));
2494 assert!(response.result.is_some());
2495 server.stop().await.unwrap();
2496 }
2497
2498 #[tokio::test]
2499 async fn http_preserves_policy_errors_and_request_ids() {
2500 use crate::security::{Permission, RbacAuthorizer, RequestPolicy};
2501 use crate::server::McpServer;
2502
2503 let reserved = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
2504 let address = reserved.local_addr().unwrap();
2505 drop(reserved);
2506
2507 let policy = RequestPolicy::new(RbacAuthorizer::new([Permission::new(
2508 "operator",
2509 methods::PING,
2510 )]));
2511 let mut server = McpServer::create("http-policy", "1.0.0").with_request_policy(policy);
2512 server
2513 .start(HttpServerTransport::new(address.to_string()))
2514 .await
2515 .unwrap();
2516
2517 let mut client = HttpClientTransport::new(format!("http://{address}"), None)
2518 .await
2519 .unwrap();
2520 let error = client
2521 .send_request(
2522 JsonRpcRequest::new(Value::from(43), methods::PING.to_string(), None::<Value>)
2523 .unwrap(),
2524 )
2525 .await
2526 .unwrap_err();
2527 assert!(matches!(error, McpError::Forbidden(_)));
2528 server.stop().await.unwrap();
2529 }
2530
2531 #[cfg(feature = "tls")]
2532 #[tokio::test]
2533 async fn mtls_server_rejects_empty_identity_before_starting() {
2534 let mut transport = HttpServerTransport::new("127.0.0.1:0")
2535 .with_mtls(MtlsServerConfig::new(Vec::new(), Vec::new(), Vec::new()));
2536 assert!(transport.start().await.is_err());
2537 assert!(!transport.is_running());
2538 }
2539}