1use std::collections::HashMap;
71use std::sync::Arc;
72
73use axum::{
74 Router,
75 extract::{
76 State, WebSocketUpgrade,
77 ws::{Message, WebSocket},
78 },
79 response::Response,
80 routing::get,
81};
82use futures::{SinkExt, StreamExt};
83use tokio::sync::{Mutex, RwLock, watch};
84
85use crate::context::{
86 ChannelClientRequester, ClientRequesterHandle, OutgoingRequest, OutgoingRequestReceiver,
87 OutgoingRequestSender, outgoing_request_channel,
88};
89use crate::error::{Error, JsonRpcError, Result};
90use crate::jsonrpc::JsonRpcService;
91use crate::protocol::{
92 JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, McpNotification,
93 RequestId,
94};
95use crate::router::{McpRouter, RouterRequest, RouterResponse};
96use crate::transport::service::{
97 CatchError, InjectAnnotations, McpBoxService, ServiceFactory, identity_factory,
98};
99use crate::{ProtocolSupport, ProtocolSupportError};
100
101struct Session {
103 id: String,
104 router: McpRouter,
105 service_factory: ServiceFactory,
106 cancel_tx: Mutex<watch::Sender<bool>>,
109}
110
111impl Session {
112 fn new(router: McpRouter, service_factory: ServiceFactory) -> Self {
113 let (cancel_tx, _) = watch::channel(false);
114 Self {
115 id: uuid::Uuid::new_v4().to_string(),
116 router,
117 service_factory,
118 cancel_tx: Mutex::new(cancel_tx),
119 }
120 }
121
122 fn make_service(&self) -> McpBoxService {
124 (self.service_factory)(self.router.clone())
125 }
126
127 async fn cancel_receiver(&self) -> watch::Receiver<bool> {
129 self.cancel_tx.lock().await.subscribe()
130 }
131
132 async fn replace_connection(&self) -> watch::Receiver<bool> {
135 let mut tx = self.cancel_tx.lock().await;
136 let _ = tx.send(true);
138 let (new_tx, new_rx) = watch::channel(false);
140 *tx = new_tx;
141 new_rx
142 }
143}
144
145impl std::fmt::Debug for Session {
146 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
147 f.debug_struct("Session")
148 .field("id", &self.id)
149 .field("router", &self.router)
150 .finish_non_exhaustive()
151 }
152}
153
154#[derive(Debug, Default)]
156struct SessionStore {
157 sessions: RwLock<HashMap<String, Arc<Session>>>,
158}
159
160impl SessionStore {
161 fn new() -> Self {
162 Self::default()
163 }
164
165 async fn create(
166 &self,
167 router: McpRouter,
168 service_factory: ServiceFactory,
169 ) -> (Arc<Session>, watch::Receiver<bool>) {
170 let session = Arc::new(Session::new(router, service_factory));
171 let cancel_rx = session.cancel_receiver().await;
172 let mut sessions = self.sessions.write().await;
173 sessions.insert(session.id.clone(), session.clone());
174 tracing::debug!(session_id = %session.id, "Created WebSocket session");
175 (session, cancel_rx)
176 }
177
178 #[cfg_attr(not(test), allow(dead_code))]
183 async fn reconnect(&self, id: &str) -> Option<(Arc<Session>, watch::Receiver<bool>)> {
184 let sessions = self.sessions.read().await;
185 let session = sessions.get(id)?;
186 let cancel_rx = session.replace_connection().await;
187 tracing::info!(session_id = %id, "Replaced active WebSocket connection (zombie prevention)");
188 Some((session.clone(), cancel_rx))
189 }
190
191 async fn remove(&self, id: &str) -> bool {
192 let mut sessions = self.sessions.write().await;
193 let removed = sessions.remove(id).is_some();
194 if removed {
195 tracing::debug!(session_id = %id, "Removed WebSocket session");
196 }
197 removed
198 }
199}
200
201struct PendingRequest {
203 response_tx: tokio::sync::oneshot::Sender<Result<serde_json::Value>>,
204}
205
206struct AppState {
208 router_template: McpRouter,
209 service_factory: ServiceFactory,
210 sessions: SessionStore,
211 protocol_support: ProtocolSupport,
212 sampling_enabled: bool,
214}
215
216pub struct WebSocketTransport {
233 router: McpRouter,
234 sampling_enabled: bool,
235 service_factory: ServiceFactory,
236 protocol_support: ProtocolSupport,
237 #[cfg(feature = "oauth")]
238 oauth_config: Option<crate::oauth::ProtectedResourceMetadata>,
239}
240
241impl WebSocketTransport {
242 pub fn new(router: McpRouter) -> Self {
244 Self {
245 router,
246 sampling_enabled: false,
247 service_factory: identity_factory(),
248 protocol_support: ProtocolSupport::default(),
249 #[cfg(feature = "oauth")]
250 oauth_config: None,
251 }
252 }
253
254 pub fn with_sampling(mut self) -> Self {
259 self.sampling_enabled = true;
260 self
261 }
262
263 pub fn protocol_support(mut self, support: ProtocolSupport) -> Self {
265 self.protocol_support = support;
266 self
267 }
268
269 pub fn protocol_versions<I, V>(
271 mut self,
272 versions: I,
273 ) -> std::result::Result<Self, ProtocolSupportError>
274 where
275 I: IntoIterator<Item = V>,
276 V: Into<String>,
277 {
278 self.protocol_support = ProtocolSupport::try_new(versions)?;
279 Ok(self)
280 }
281
282 #[cfg(feature = "oauth")]
303 pub fn oauth(mut self, metadata: crate::oauth::ProtectedResourceMetadata) -> Self {
304 self.oauth_config = Some(metadata);
305 self
306 }
307
308 #[cfg(feature = "oauth")]
314 pub fn into_oauth_router<V>(
315 self,
316 validator: V,
317 metadata: crate::oauth::ProtectedResourceMetadata,
318 policy: crate::oauth::ScopePolicy,
319 ) -> std::result::Result<Router, crate::oauth::ProtectedResourceMetadataError>
320 where
321 V: crate::oauth::TokenValidator,
322 {
323 metadata.validate()?;
324 let oauth_layer =
325 crate::oauth::OAuthLayer::new(validator, metadata.clone()).scope_policy(policy.clone());
326 let router = self
327 .layer(crate::oauth::ScopeEnforcementLayer::new(policy))
328 .oauth(metadata)
329 .into_router();
330 Ok(router.layer(oauth_layer))
331 }
332
333 #[cfg(feature = "oauth")]
335 pub fn into_oauth_router_at<V>(
336 self,
337 path: &str,
338 validator: V,
339 metadata: crate::oauth::ProtectedResourceMetadata,
340 policy: crate::oauth::ScopePolicy,
341 ) -> std::result::Result<Router, crate::oauth::ProtectedResourceMetadataError>
342 where
343 V: crate::oauth::TokenValidator,
344 {
345 metadata.validate()?;
346 let oauth_layer =
347 crate::oauth::OAuthLayer::new(validator, metadata.clone()).scope_policy(policy.clone());
348 let router = self
349 .layer(crate::oauth::ScopeEnforcementLayer::new(policy))
350 .oauth(metadata)
351 .into_router_at(path);
352 Ok(router.layer(oauth_layer))
353 }
354
355 pub fn layer<L>(mut self, layer: L) -> Self
384 where
385 L: tower::Layer<McpRouter> + Send + Sync + 'static,
386 L::Service:
387 tower::Service<RouterRequest, Response = RouterResponse> + Clone + Send + 'static,
388 <L::Service as tower::Service<RouterRequest>>::Error: std::fmt::Display + Send,
389 <L::Service as tower::Service<RouterRequest>>::Future: Send,
390 {
391 self.service_factory = Arc::new(move |router: McpRouter| {
392 let annotations = router.tool_annotations_map();
393 let wrapped = layer.layer(router);
394 tower::util::BoxCloneService::new(InjectAnnotations::new(
395 CatchError::new(wrapped),
396 annotations,
397 ))
398 });
399 self
400 }
401
402 pub fn into_router(self) -> Router {
404 #[cfg(feature = "oauth")]
405 let oauth_config = self.oauth_config;
406
407 let state = Arc::new(AppState {
408 router_template: self.router,
409 service_factory: self.service_factory,
410 sessions: SessionStore::new(),
411 protocol_support: self.protocol_support,
412 sampling_enabled: self.sampling_enabled,
413 });
414
415 let router = Router::new()
416 .route("/", get(handle_websocket))
417 .with_state(state);
418
419 #[cfg(feature = "oauth")]
420 let router = add_oauth_route(router, "", oauth_config.as_ref());
421
422 router
423 }
424
425 pub fn into_router_at(self, path: &str) -> Router {
427 #[cfg(feature = "oauth")]
428 let oauth_config = self.oauth_config;
429
430 let state = Arc::new(AppState {
431 router_template: self.router,
432 service_factory: self.service_factory,
433 sessions: SessionStore::new(),
434 protocol_support: self.protocol_support,
435 sampling_enabled: self.sampling_enabled,
436 });
437
438 let ws_router = Router::new()
439 .route("/", get(handle_websocket))
440 .with_state(state);
441
442 let router = Router::new().nest(path, ws_router);
443
444 #[cfg(feature = "oauth")]
445 let router = add_oauth_route(router, path, oauth_config.as_ref());
446
447 router
448 }
449
450 pub async fn serve(self, addr: &str) -> Result<()> {
452 let listener = tokio::net::TcpListener::bind(addr)
453 .await
454 .map_err(|e| Error::Transport(format!("Failed to bind to {}: {}", addr, e)))?;
455
456 tracing::info!("MCP WebSocket transport listening on {}", addr);
457
458 let router = self.into_router();
459 axum::serve(listener, router)
460 .await
461 .map_err(|e| Error::Transport(format!("Server error: {}", e)))?;
462
463 Ok(())
464 }
465}
466
467#[cfg(feature = "oauth")]
469fn add_oauth_route(
470 router: Router,
471 _base_path: &str,
472 metadata: Option<&crate::oauth::ProtectedResourceMetadata>,
473) -> Router {
474 if let Some(metadata) = metadata {
475 let metadata = metadata.clone();
476 let well_known_path =
477 crate::oauth::ProtectedResourceMetadata::well_known_path_for_resource(
478 &metadata.resource,
479 )
480 .unwrap_or_else(|_| {
481 crate::oauth::ProtectedResourceMetadata::well_known_path().to_string()
482 });
483 router.route(
484 &well_known_path,
485 get(move || {
486 let m = metadata.clone();
487 async move { axum::Json(m) }
488 }),
489 )
490 } else {
491 router
492 }
493}
494
495#[derive(Debug, Default)]
500struct McpSubprotocols {
501 auth_token: Option<String>,
503 protocol_version: Option<String>,
505 selected: Vec<String>,
507}
508
509fn parse_mcp_subprotocols(
513 headers: &axum::http::HeaderMap,
514 protocol_support: &ProtocolSupport,
515) -> McpSubprotocols {
516 let mut result = McpSubprotocols::default();
517
518 let Some(header) = headers.get("sec-websocket-protocol") else {
519 return result;
520 };
521 let Ok(header_str) = header.to_str() else {
522 return result;
523 };
524
525 for protocol in header_str.split(',').map(|s| s.trim()) {
526 if let Some(token) = protocol.strip_prefix("mcp.auth.") {
527 if !token.is_empty() {
528 result.auth_token = Some(token.to_string());
529 result.selected.push(protocol.to_string());
530 }
531 } else if let Some(version) = protocol.strip_prefix("mcp.version.") {
532 if protocol_support.contains(version) {
533 result.protocol_version = Some(version.to_string());
534 result.selected.push(protocol.to_string());
535 } else {
536 tracing::warn!(version = %version, "Unsupported MCP protocol version in subprotocol");
537 }
538 }
539 }
540
541 result
542}
543
544async fn handle_websocket(
550 State(state): State<Arc<AppState>>,
551 request: axum::extract::Request,
552) -> Response {
553 use axum::extract::FromRequestParts;
554 use axum::response::IntoResponse;
555
556 let (mut parts, _body) = request.into_parts();
557
558 let subprotocols = parse_mcp_subprotocols(&parts.headers, &state.protocol_support);
560 if let Some(ref version) = subprotocols.protocol_version {
561 tracing::debug!(version = %version, "Client requested MCP protocol version via subprotocol");
562 }
563
564 #[allow(unused_mut)]
566 let mut mcp_extensions = crate::router::Extensions::new();
567 #[cfg(feature = "oauth")]
568 {
569 if let Some(claims) = parts.extensions.get::<crate::oauth::token::TokenClaims>() {
570 mcp_extensions.insert(claims.clone());
571 }
572 }
573
574 if let Some(ref token) = subprotocols.auth_token {
576 mcp_extensions.insert(WebSocketAuthToken(token.clone()));
577 }
578 if let Some(ref version) = subprotocols.protocol_version
579 && let Ok(revision) = version.parse::<crate::inspection::McpProtocolRevision>()
580 {
581 mcp_extensions.insert(revision);
582 }
583
584 let ws: WebSocketUpgrade = match WebSocketUpgrade::from_request_parts(&mut parts, &()).await {
586 Ok(ws) => ws,
587 Err(e) => return e.into_response(),
588 };
589
590 let ws = if !subprotocols.selected.is_empty() {
592 ws.protocols(subprotocols.selected)
593 } else {
594 ws
595 };
596
597 ws.on_upgrade(move |socket| handle_socket(socket, state, mcp_extensions))
598}
599
600#[derive(Debug, Clone)]
605pub struct WebSocketAuthToken(pub String);
606
607async fn handle_socket(
609 socket: WebSocket,
610 state: Arc<AppState>,
611 mcp_extensions: crate::router::Extensions,
612) {
613 let (session, cancel_rx) = state
615 .sessions
616 .create(
617 state.router_template.with_fresh_session(),
618 state.service_factory.clone(),
619 )
620 .await;
621 let session_id = session.id.clone();
622 let protocol_support = state.protocol_support.clone();
623
624 tracing::info!(session_id = %session_id, "WebSocket connection established");
625
626 if state.sampling_enabled {
627 handle_socket_bidirectional(
628 socket,
629 session,
630 &session_id,
631 mcp_extensions,
632 protocol_support,
633 cancel_rx,
634 )
635 .await;
636 } else {
637 handle_socket_simple(
638 socket,
639 session,
640 &session_id,
641 mcp_extensions,
642 protocol_support,
643 cancel_rx,
644 )
645 .await;
646 }
647
648 state.sessions.remove(&session_id).await;
650 tracing::info!(session_id = %session_id, "WebSocket connection closed");
651}
652
653async fn handle_socket_simple(
655 socket: WebSocket,
656 session: Arc<Session>,
657 session_id: &str,
658 mcp_extensions: crate::router::Extensions,
659 protocol_support: ProtocolSupport,
660 mut cancel_rx: watch::Receiver<bool>,
661) {
662 let mut service = JsonRpcService::new(session.make_service())
663 .with_extensions(mcp_extensions)
664 .protocol_support(protocol_support);
665 let (mut sender, mut receiver) = socket.split();
666
667 loop {
669 let msg = tokio::select! {
670 msg = receiver.next() => {
671 match msg {
672 Some(msg) => msg,
673 None => break,
674 }
675 }
676 _ = cancel_rx.changed() => {
677 if *cancel_rx.borrow() {
678 tracing::info!(session_id = %session_id, "Connection superseded by new connection, closing");
679 let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
680 code: 1000,
681 reason: "Connection replaced by newer WebSocket connection".into(),
682 }))).await;
683 break;
684 }
685 continue;
686 }
687 };
688 let msg = match msg {
689 Ok(m) => m,
690 Err(e) => {
691 tracing::error!(error = %e, "WebSocket receive error");
692 break;
693 }
694 };
695
696 match msg {
697 Message::Text(text) => {
698 match process_message(&mut service, &session.router, &text).await {
699 Ok(Some(response)) => {
700 let response_json = match serde_json::to_string(&response) {
701 Ok(json) => json,
702 Err(e) => {
703 tracing::error!(error = %e, "Failed to serialize response");
704 continue;
705 }
706 };
707
708 if let Err(e) = sender.send(Message::Text(response_json.into())).await {
709 tracing::error!(error = %e, "Failed to send response");
710 break;
711 }
712 }
713 Ok(None) => {
714 }
716 Err(e) => {
717 tracing::error!(error = %e, "Error processing message");
718 let error_response = JsonRpcResponse::error(
719 None,
720 JsonRpcError::internal_error(e.to_string()),
721 );
722 if let Ok(json) = serde_json::to_string(&error_response) {
723 let _ = sender.send(Message::Text(json.into())).await;
724 }
725 }
726 }
727 }
728 Message::Binary(_) => {
729 tracing::warn!(session_id = %session_id, "Received binary frame, closing with 1003");
732 let _ = sender
733 .send(Message::Close(Some(axum::extract::ws::CloseFrame {
734 code: 1003,
735 reason: "Binary frames are not supported by MCP".into(),
736 })))
737 .await;
738 break;
739 }
740 Message::Ping(data) => {
741 if let Err(e) = sender.send(Message::Pong(data)).await {
742 tracing::error!(error = %e, "Failed to send pong");
743 break;
744 }
745 }
746 Message::Pong(_) => {
747 }
749 Message::Close(_) => {
750 tracing::info!(session_id = %session_id, "WebSocket close received");
751 break;
752 }
753 }
754 }
755}
756
757async fn handle_socket_bidirectional(
759 socket: WebSocket,
760 session: Arc<Session>,
761 session_id: &str,
762 _mcp_extensions: crate::router::Extensions,
763 protocol_support: ProtocolSupport,
764 mut cancel_rx: watch::Receiver<bool>,
765) {
766 let (request_tx, mut request_rx): (OutgoingRequestSender, OutgoingRequestReceiver) =
768 outgoing_request_channel(32);
769
770 let client_requester: ClientRequesterHandle = Arc::new(ChannelClientRequester::new(request_tx));
772
773 let router = session
775 .router
776 .clone()
777 .with_client_requester(client_requester);
778 let mut service = JsonRpcService::new((session.service_factory)(router.clone()))
779 .with_extensions(_mcp_extensions)
780 .protocol_support(protocol_support);
781
782 let pending_requests: Arc<Mutex<HashMap<RequestId, PendingRequest>>> =
784 Arc::new(Mutex::new(HashMap::new()));
785
786 let (sender, mut receiver) = socket.split();
787 let sender = Arc::new(Mutex::new(sender));
788
789 let session_id_owned = session_id.to_string();
790
791 loop {
792 tokio::select! {
793 msg = receiver.next() => {
795 let msg = match msg {
796 Some(Ok(m)) => m,
797 Some(Err(e)) => {
798 tracing::error!(error = %e, "WebSocket receive error");
799 break;
800 }
801 None => break,
802 };
803
804 match msg {
805 Message::Text(text) => {
806 let result = handle_incoming_message(
807 &text,
808 &mut service,
809 &router,
810 pending_requests.clone(),
811 sender.clone(),
812 ).await;
813 if let Err(e) = result {
814 tracing::error!(error = %e, "Error handling incoming message");
815 }
816 }
817 Message::Binary(_) => {
818 tracing::warn!(session_id = %session_id_owned, "Received binary frame, closing with 1003");
821 let mut s = sender.lock().await;
822 let _ = s.send(Message::Close(Some(axum::extract::ws::CloseFrame {
823 code: 1003,
824 reason: "Binary frames are not supported by MCP".into(),
825 }))).await;
826 break;
827 }
828 Message::Ping(data) => {
829 let mut sender = sender.lock().await;
830 if let Err(e) = sender.send(Message::Pong(data)).await {
831 tracing::error!(error = %e, "Failed to send pong");
832 break;
833 }
834 }
835 Message::Pong(_) => {}
836 Message::Close(_) => {
837 tracing::info!(session_id = %session_id_owned, "WebSocket close received");
838 break;
839 }
840 }
841 }
842
843 Some(outgoing) = request_rx.recv() => {
845 let result = send_outgoing_request(
846 outgoing,
847 pending_requests.clone(),
848 sender.clone(),
849 ).await;
850 if let Err(e) = result {
851 tracing::error!(error = %e, "Error sending outgoing request");
852 }
853 }
854
855 _ = cancel_rx.changed() => {
857 if *cancel_rx.borrow() {
858 tracing::info!(session_id = %session_id_owned, "Connection superseded by new connection, closing");
859 let mut s = sender.lock().await;
860 let _ = s.send(Message::Close(Some(axum::extract::ws::CloseFrame {
861 code: 1000,
862 reason: "Connection replaced by newer WebSocket connection".into(),
863 }))).await;
864 break;
865 }
866 }
867 }
868 }
869}
870
871async fn handle_incoming_message<S>(
873 text: &str,
874 service: &mut JsonRpcService<McpBoxService>,
875 router: &McpRouter,
876 pending_requests: Arc<Mutex<HashMap<RequestId, PendingRequest>>>,
877 sender: Arc<Mutex<S>>,
878) -> Result<()>
879where
880 S: futures::Sink<Message> + Unpin,
881 S::Error: std::fmt::Display,
882{
883 let parsed: serde_json::Value = serde_json::from_str(text)?;
884
885 if let Err(error) =
886 service.inspect_incoming_value(&parsed, crate::inspection::McpDirection::ClientToServer)
887 {
888 let response = JsonRpcResponse::error(None, error);
889 let response = serde_json::to_string(&response)
890 .map_err(|error| Error::Transport(format!("Failed to serialize response: {error}")))?;
891 sender
892 .lock()
893 .await
894 .send(Message::Text(response.into()))
895 .await
896 .map_err(|error| Error::Transport(format!("Failed to send response: {error}")))?;
897 return Ok(());
898 }
899
900 if parsed.get("method").is_none()
902 && (parsed.get("result").is_some() || parsed.get("error").is_some())
903 {
904 return handle_response(&parsed, pending_requests).await;
905 }
906
907 if !parsed.is_array() && parsed.get("id").is_none() {
909 if let Ok(notification) = serde_json::from_str::<JsonRpcNotification>(text) {
910 let mcp_notification = McpNotification::from_jsonrpc(¬ification)?;
911 router.handle_notification(mcp_notification);
912 }
913 return Ok(());
914 }
915
916 let message: JsonRpcMessage = serde_json::from_str(text)?;
918 match service.call_message(message).await {
919 Ok(response) => {
920 let response_json = serde_json::to_string(&response)
921 .map_err(|e| Error::Transport(format!("Failed to serialize response: {}", e)))?;
922 let mut sender = sender.lock().await;
923 sender
924 .send(Message::Text(response_json.into()))
925 .await
926 .map_err(|e| Error::Transport(format!("Failed to send response: {}", e)))?;
927 }
928 Err(e) => {
929 tracing::error!(error = %e, "Error processing message");
930 let error_response =
931 JsonRpcResponse::error(None, JsonRpcError::internal_error(e.to_string()));
932 if let Ok(json) = serde_json::to_string(&error_response) {
933 let mut sender = sender.lock().await;
934 let _ = sender.send(Message::Text(json.into())).await;
935 }
936 }
937 }
938
939 Ok(())
940}
941
942async fn handle_response(
944 parsed: &serde_json::Value,
945 pending_requests: Arc<Mutex<HashMap<RequestId, PendingRequest>>>,
946) -> Result<()> {
947 let id = match parsed.get("id") {
948 Some(id) => {
949 if let Some(n) = id.as_i64() {
950 RequestId::Number(n)
951 } else if let Some(s) = id.as_str() {
952 RequestId::String(s.to_string())
953 } else {
954 tracing::warn!("Response has invalid id type");
955 return Ok(());
956 }
957 }
958 None => {
959 tracing::warn!("Response missing id field");
960 return Ok(());
961 }
962 };
963
964 let pending = {
965 let mut pending_requests = pending_requests.lock().await;
966 pending_requests.remove(&id)
967 };
968
969 match pending {
970 Some(pending) => {
971 let result = if let Some(error) = parsed.get("error") {
972 let code = error.get("code").and_then(|c| c.as_i64()).unwrap_or(-1);
973 let message = error
974 .get("message")
975 .and_then(|m| m.as_str())
976 .unwrap_or("Unknown error");
977 Err(Error::Internal(format!(
978 "Client error ({}): {}",
979 code, message
980 )))
981 } else if let Some(result) = parsed.get("result") {
982 Ok(result.clone())
983 } else {
984 Err(Error::Internal(
985 "Response has neither result nor error".to_string(),
986 ))
987 };
988
989 let _ = pending.response_tx.send(result);
991 }
992 None => {
993 tracing::warn!(id = ?id, "Received response for unknown request");
994 }
995 }
996
997 Ok(())
998}
999
1000async fn send_outgoing_request<S>(
1002 outgoing: OutgoingRequest,
1003 pending_requests: Arc<Mutex<HashMap<RequestId, PendingRequest>>>,
1004 sender: Arc<Mutex<S>>,
1005) -> Result<()>
1006where
1007 S: futures::Sink<Message> + Unpin,
1008 S::Error: std::fmt::Display,
1009{
1010 let request = JsonRpcRequest {
1012 jsonrpc: "2.0".to_string(),
1013 id: outgoing.id.clone(),
1014 method: outgoing.method,
1015 params: Some(outgoing.params),
1016 };
1017
1018 let request_json = serde_json::to_string(&request)
1019 .map_err(|e| Error::Transport(format!("Failed to serialize request: {}", e)))?;
1020
1021 tracing::debug!(output = %request_json, "Sending request to client");
1022
1023 {
1025 let mut pending = pending_requests.lock().await;
1026 pending.insert(
1027 outgoing.id,
1028 PendingRequest {
1029 response_tx: outgoing.response_tx,
1030 },
1031 );
1032 }
1033
1034 let mut sender = sender.lock().await;
1036 sender
1037 .send(Message::Text(request_json.into()))
1038 .await
1039 .map_err(|e| Error::Transport(format!("Failed to send request: {}", e)))?;
1040
1041 Ok(())
1042}
1043
1044async fn process_message(
1046 service: &mut JsonRpcService<McpBoxService>,
1047 router: &McpRouter,
1048 text: &str,
1049) -> Result<Option<crate::protocol::JsonRpcResponseMessage>> {
1050 let parsed: serde_json::Value = serde_json::from_str(text)?;
1052 if let Err(error) =
1053 service.inspect_incoming_value(&parsed, crate::inspection::McpDirection::ClientToServer)
1054 {
1055 return Ok(Some(crate::protocol::JsonRpcResponseMessage::Single(
1056 JsonRpcResponse::error(None, error),
1057 )));
1058 }
1059 if !parsed.is_array()
1060 && parsed.get("id").is_none()
1061 && let Ok(notification) = serde_json::from_str::<JsonRpcNotification>(text)
1062 {
1063 let mcp_notification = McpNotification::from_jsonrpc(¬ification)?;
1064 router.handle_notification(mcp_notification);
1065 return Ok(None);
1066 }
1067
1068 let message: JsonRpcMessage = serde_json::from_str(text)?;
1070 let response = service.call_message(message).await?;
1071 Ok(Some(response))
1072}
1073
1074#[cfg(test)]
1075mod tests {
1076 use super::*;
1077
1078 fn create_test_router() -> McpRouter {
1079 McpRouter::new().server_info("test-server", "1.0.0")
1080 }
1081
1082 #[tokio::test]
1083 async fn test_websocket_transport_builds() {
1084 let transport = WebSocketTransport::new(create_test_router());
1085 let _router = transport.into_router();
1086 }
1087
1088 #[tokio::test]
1089 async fn test_websocket_transport_at_path() {
1090 let transport = WebSocketTransport::new(create_test_router());
1091 let _router = transport.into_router_at("/mcp");
1092 }
1093
1094 #[cfg(feature = "oauth")]
1095 #[tokio::test]
1096 async fn test_oauth_metadata_route_is_path_aware() {
1097 use axum::body::Body;
1098 use axum::http::{Request, StatusCode};
1099 use tower::ServiceExt;
1100
1101 let metadata =
1102 crate::oauth::ProtectedResourceMetadata::new("https://mcp.example.com/tenant/ws")
1103 .authorization_server("https://auth.example.com");
1104 let app = WebSocketTransport::new(create_test_router())
1105 .oauth(metadata)
1106 .into_router_at("/tenant/ws");
1107 let request = Request::builder()
1108 .uri("/.well-known/oauth-protected-resource/tenant/ws")
1109 .body(Body::empty())
1110 .unwrap();
1111 let response = app.oneshot(request).await.unwrap();
1112
1113 assert_eq!(response.status(), StatusCode::OK);
1114 }
1115
1116 #[tokio::test]
1117 async fn test_layer_with_identity() {
1118 let transport = WebSocketTransport::new(create_test_router())
1120 .layer(tower::layer::util::Identity::new());
1121 let _router = transport.into_router();
1122 }
1123
1124 #[tokio::test]
1125 async fn test_layer_with_timeout() {
1126 use std::time::Duration;
1127 use tower::timeout::TimeoutLayer;
1128
1129 let transport = WebSocketTransport::new(create_test_router())
1130 .layer(TimeoutLayer::new(Duration::from_secs(30)));
1131 let _router = transport.into_router();
1132 }
1133
1134 #[tokio::test]
1135 async fn test_layer_with_composed_layers() {
1136 use std::time::Duration;
1137 use tower::ServiceBuilder;
1138 use tower::timeout::TimeoutLayer;
1139
1140 let transport = WebSocketTransport::new(create_test_router()).layer(
1141 ServiceBuilder::new()
1142 .layer(TimeoutLayer::new(Duration::from_secs(30)))
1143 .concurrency_limit(100)
1144 .into_inner(),
1145 );
1146 let _router = transport.into_router();
1147 }
1148
1149 #[test]
1150 fn test_parse_mcp_subprotocols_empty() {
1151 let headers = axum::http::HeaderMap::new();
1152 let result = parse_mcp_subprotocols(&headers, &ProtocolSupport::default());
1153 assert!(result.auth_token.is_none());
1154 assert!(result.protocol_version.is_none());
1155 assert!(result.selected.is_empty());
1156 }
1157
1158 #[test]
1159 fn test_parse_mcp_subprotocols_auth_and_version() {
1160 let mut headers = axum::http::HeaderMap::new();
1161 headers.insert(
1162 "sec-websocket-protocol",
1163 "mcp.auth.my-secret-token, mcp.version.2025-11-25"
1164 .parse()
1165 .unwrap(),
1166 );
1167 let result = parse_mcp_subprotocols(&headers, &ProtocolSupport::default());
1168 assert_eq!(result.auth_token.as_deref(), Some("my-secret-token"));
1169 assert_eq!(result.protocol_version.as_deref(), Some("2025-11-25"));
1170 assert_eq!(result.selected.len(), 2);
1171 }
1172
1173 #[test]
1174 fn test_parse_mcp_subprotocols_unsupported_version() {
1175 let mut headers = axum::http::HeaderMap::new();
1176 headers.insert(
1177 "sec-websocket-protocol",
1178 "mcp.version.1999-01-01".parse().unwrap(),
1179 );
1180 let result = parse_mcp_subprotocols(&headers, &ProtocolSupport::default());
1181 assert!(result.protocol_version.is_none());
1182 assert!(result.selected.is_empty());
1183 }
1184
1185 #[test]
1186 fn test_parse_mcp_subprotocols_older_supported_version() {
1187 let mut headers = axum::http::HeaderMap::new();
1188 headers.insert(
1189 "sec-websocket-protocol",
1190 "mcp.version.2025-03-26".parse().unwrap(),
1191 );
1192 let result = parse_mcp_subprotocols(&headers, &ProtocolSupport::default());
1193 assert_eq!(result.protocol_version.as_deref(), Some("2025-03-26"));
1194 assert_eq!(result.selected.len(), 1);
1195 }
1196
1197 #[test]
1198 fn test_parse_mcp_subprotocols_auth_only() {
1199 let mut headers = axum::http::HeaderMap::new();
1200 headers.insert(
1201 "sec-websocket-protocol",
1202 "mcp.auth.bearer-xyz123".parse().unwrap(),
1203 );
1204 let result = parse_mcp_subprotocols(&headers, &ProtocolSupport::default());
1205 assert_eq!(result.auth_token.as_deref(), Some("bearer-xyz123"));
1206 assert!(result.protocol_version.is_none());
1207 }
1208
1209 #[test]
1210 fn test_parse_mcp_subprotocols_ignores_unknown() {
1211 let mut headers = axum::http::HeaderMap::new();
1212 headers.insert(
1213 "sec-websocket-protocol",
1214 "graphql-ws, mcp.auth.token, mcp.version.2025-11-25, other-protocol"
1215 .parse()
1216 .unwrap(),
1217 );
1218 let result = parse_mcp_subprotocols(&headers, &ProtocolSupport::default());
1219 assert_eq!(result.auth_token.as_deref(), Some("token"));
1220 assert_eq!(result.protocol_version.as_deref(), Some("2025-11-25"));
1221 assert_eq!(result.selected.len(), 2);
1223 }
1224
1225 #[cfg(feature = "stateless")]
1226 #[test]
1227 fn websocket_subprotocol_uses_exact_runtime_allow_list() {
1228 let mut headers = axum::http::HeaderMap::new();
1229 headers.insert(
1230 "sec-websocket-protocol",
1231 "mcp.version.2026-07-28, mcp.version.2025-11-25"
1232 .parse()
1233 .unwrap(),
1234 );
1235
1236 let stable = parse_mcp_subprotocols(&headers, &ProtocolSupport::stable());
1237 assert_eq!(stable.protocol_version.as_deref(), Some("2025-11-25"));
1238 assert_eq!(stable.selected, vec!["mcp.version.2025-11-25"]);
1239
1240 let final_only = ProtocolSupport::try_new(["2026-07-28"]).unwrap();
1241 let final_selected = parse_mcp_subprotocols(&headers, &final_only);
1242 assert_eq!(
1243 final_selected.protocol_version.as_deref(),
1244 Some("2026-07-28")
1245 );
1246 assert_eq!(final_selected.selected, vec!["mcp.version.2026-07-28"]);
1247 }
1248
1249 #[cfg(feature = "stateless")]
1250 #[tokio::test]
1251 async fn websocket_request_body_selects_final_lifecycle() {
1252 let router = create_test_router();
1253 let service = identity_factory()(router.clone());
1254 let mut service = JsonRpcService::new(service)
1255 .protocol_support(ProtocolSupport::try_new(["2026-07-28"]).unwrap());
1256 let request = serde_json::json!({
1257 "jsonrpc": "2.0",
1258 "id": 1,
1259 "method": "server/discover",
1260 "params": {
1261 "_meta": {
1262 "io.modelcontextprotocol/protocolVersion": "2026-07-28",
1263 "io.modelcontextprotocol/clientCapabilities": {}
1264 }
1265 }
1266 });
1267
1268 let response = process_message(&mut service, &router, &request.to_string())
1269 .await
1270 .unwrap()
1271 .expect("request response");
1272 let response = serde_json::to_value(response).unwrap();
1273 assert_eq!(response["result"]["resultType"], "complete");
1274 assert_eq!(response["result"]["ttlMs"], 0);
1275 assert_eq!(response["result"]["cacheScope"], "private");
1276 assert_eq!(response["result"]["supportedVersions"][0], "2026-07-28");
1277 }
1278
1279 async fn websocket_batch_response(revision: &str) -> serde_json::Value {
1280 let router = create_test_router();
1281 let service = identity_factory()(router.clone());
1282 let mut service = JsonRpcService::new(service);
1283 let initialize = serde_json::json!({
1284 "jsonrpc": "2.0",
1285 "id": 0,
1286 "method": "initialize",
1287 "params": {
1288 "protocolVersion": revision,
1289 "capabilities": {},
1290 "clientInfo": {"name": "test", "version": "1.0"}
1291 }
1292 });
1293 process_message(&mut service, &router, &initialize.to_string())
1294 .await
1295 .unwrap();
1296 router.handle_notification(McpNotification::Initialized);
1297
1298 let batch = serde_json::json!([
1299 {"jsonrpc": "2.0", "id": 1, "method": "ping"},
1300 {"jsonrpc": "2.0", "id": 2, "method": "tools/list"}
1301 ]);
1302 let response = process_message(&mut service, &router, &batch.to_string())
1303 .await
1304 .unwrap()
1305 .expect("batch response");
1306 serde_json::to_value(response).unwrap()
1307 }
1308
1309 #[tokio::test]
1310 async fn websocket_accepts_batch_for_2025_03() {
1311 let response = websocket_batch_response("2025-03-26").await;
1312 assert_eq!(response.as_array().map(Vec::len), Some(2));
1313 }
1314
1315 #[tokio::test]
1316 async fn websocket_rejects_batch_for_2025_11() {
1317 let response = websocket_batch_response("2025-11-25").await;
1318 assert_eq!(response["error"]["code"], -32600);
1319 assert!(
1320 response["error"]["message"]
1321 .as_str()
1322 .unwrap()
1323 .contains("does not permit top-level JSON-RPC batches")
1324 );
1325 }
1326
1327 #[tokio::test]
1328 async fn test_session_cancel_receiver() {
1329 let router = create_test_router();
1330 let session = Session::new(router, identity_factory());
1331 let mut rx = session.cancel_receiver().await;
1332
1333 assert!(!*rx.borrow());
1335
1336 let _new_rx = session.replace_connection().await;
1338 rx.changed().await.unwrap();
1339 assert!(*rx.borrow());
1340 }
1341
1342 #[tokio::test]
1343 async fn test_session_replace_connection_new_rx_starts_clean() {
1344 let router = create_test_router();
1345 let session = Session::new(router, identity_factory());
1346
1347 let _rx1 = session.cancel_receiver().await;
1349
1350 let rx2 = session.replace_connection().await;
1352 assert!(!*rx2.borrow(), "New receiver should start as not-cancelled");
1353 }
1354
1355 #[tokio::test]
1356 async fn test_session_store_reconnect() {
1357 let router = create_test_router();
1358 let store = SessionStore::new();
1359
1360 let (session, mut rx1) = store
1361 .create(router.with_fresh_session(), identity_factory())
1362 .await;
1363 let session_id = session.id.clone();
1364
1365 let result = store.reconnect(&session_id).await;
1367 assert!(result.is_some());
1368 let (_session2, rx2) = result.unwrap();
1369
1370 rx1.changed().await.unwrap();
1372 assert!(*rx1.borrow());
1373
1374 assert!(!*rx2.borrow());
1376 }
1377}