Skip to main content

tower_mcp/transport/
websocket.rs

1//! WebSocket transport for MCP
2//!
3//! Provides full-duplex communication over WebSocket, ideal for:
4//! - Bidirectional notifications
5//! - Long-lived connections
6//! - Lower latency than HTTP polling
7//! - Server-to-client requests (sampling)
8//!
9//! # Example
10//!
11//! ```rust,no_run
12//! use tower_mcp::{BoxError, McpRouter, ToolBuilder, CallToolResult};
13//! use tower_mcp::transport::websocket::WebSocketTransport;
14//! use schemars::JsonSchema;
15//! use serde::Deserialize;
16//!
17//! #[derive(Debug, Deserialize, JsonSchema)]
18//! struct Input { value: String }
19//!
20//! #[tokio::main]
21//! async fn main() -> Result<(), BoxError> {
22//!     let tool = ToolBuilder::new("echo")
23//!         .handler(|i: Input| async move { Ok(CallToolResult::text(i.value)) })
24//!         .build();
25//!
26//!     let router = McpRouter::new()
27//!         .server_info("my-server", "1.0.0")
28//!         .tool(tool);
29//!
30//!     let transport = WebSocketTransport::new(router);
31//!     transport.serve("127.0.0.1:3000").await?;
32//!     Ok(())
33//! }
34//! ```
35//!
36//! # Sampling Support
37//!
38//! The WebSocket transport supports server-to-client requests like sampling.
39//! Use [`WebSocketTransport::new`] with [`with_sampling`](WebSocketTransport::with_sampling) to enable:
40//!
41//! ```rust,no_run
42//! use tower_mcp::{BoxError, McpRouter, ToolBuilder, CallToolResult, CreateMessageParams, SamplingMessage};
43//! use tower_mcp::extract::{Context, RawArgs};
44//! use tower_mcp::transport::websocket::WebSocketTransport;
45//!
46//! #[tokio::main]
47//! async fn main() -> Result<(), BoxError> {
48//!     let tool = ToolBuilder::new("ai-tool")
49//!         .extractor_handler((), |ctx: Context, RawArgs(_): RawArgs| async move {
50//!             // Request LLM completion from client
51//!             let params = CreateMessageParams::new(
52//!                 vec![SamplingMessage::user("Summarize this...")],
53//!                 500,
54//!             );
55//!             let result = ctx.sample(params).await?;
56//!             Ok(CallToolResult::text(format!("{:?}", result.content)))
57//!         })
58//!         .build();
59//!
60//!     let router = McpRouter::new()
61//!         .server_info("my-server", "1.0.0")
62//!         .tool(tool);
63//!
64//!     let transport = WebSocketTransport::new(router).with_sampling();
65//!     transport.serve("127.0.0.1:3000").await?;
66//!     Ok(())
67//! }
68//! ```
69
70use 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
101/// Session state for WebSocket transport
102struct Session {
103    id: String,
104    router: McpRouter,
105    service_factory: ServiceFactory,
106    /// Sender to signal the active connection to close (zombie prevention).
107    /// Sending `true` tells the current connection to shut down.
108    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    /// Create a middleware-wrapped service from this session's router.
123    fn make_service(&self) -> McpBoxService {
124        (self.service_factory)(self.router.clone())
125    }
126
127    /// Get a receiver that will be notified when this connection should close.
128    async fn cancel_receiver(&self) -> watch::Receiver<bool> {
129        self.cancel_tx.lock().await.subscribe()
130    }
131
132    /// Signal the current active connection to close and create a fresh
133    /// cancellation channel for the replacement connection.
134    async fn replace_connection(&self) -> watch::Receiver<bool> {
135        let mut tx = self.cancel_tx.lock().await;
136        // Signal the old connection to shut down
137        let _ = tx.send(true);
138        // Replace with a fresh channel so new subscribers start clean
139        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/// Session store for WebSocket connections
155#[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    /// Look up an existing session by ID and replace its active connection.
179    ///
180    /// Signals the previous connection to close and returns a new cancellation
181    /// receiver for the replacement connection.
182    #[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
201/// Pending request waiting for a response
202struct PendingRequest {
203    response_tx: tokio::sync::oneshot::Sender<Result<serde_json::Value>>,
204}
205
206/// Shared state for WebSocket transport
207struct AppState {
208    router_template: McpRouter,
209    service_factory: ServiceFactory,
210    sessions: SessionStore,
211    protocol_support: ProtocolSupport,
212    /// Whether sampling is enabled
213    sampling_enabled: bool,
214}
215
216/// WebSocket transport for MCP servers
217///
218/// Provides full-duplex communication over WebSocket.
219///
220/// WebSocket is a tower-mcp custom transport binding, not a standard
221/// 2026-07-28 MCP transport. JSON-RPC request bodies remain the source of
222/// truth for final per-request metadata. The optional `mcp.version.*`
223/// subprotocol is an upgrade-time compatibility hint constrained by this
224/// transport's exact [`ProtocolSupport`] allow-list.
225///
226/// Connection and cancellation semantics remain those of this custom binding:
227/// a WebSocket close terminates the connection, and reconnecting a stored
228/// session closes the older socket. It does not currently cancel an in-flight
229/// handler. `notifications/cancelled` is processed between requests, so it
230/// cannot interrupt a handler that is already executing. Final
231/// `subscriptions/listen` multiplexing is not implemented on this binding.
232pub 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    /// Create a new WebSocket transport
243    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    /// Enable sampling support for this transport.
255    ///
256    /// When sampling is enabled, tool handlers can use `ctx.sample()` to
257    /// request LLM completions from connected clients.
258    pub fn with_sampling(mut self) -> Self {
259        self.sampling_enabled = true;
260        self
261    }
262
263    /// Set the exact protocol versions this custom binding accepts.
264    pub fn protocol_support(mut self, support: ProtocolSupport) -> Self {
265        self.protocol_support = support;
266        self
267    }
268
269    /// Construct and set an exact runtime protocol-version allow-list.
270    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    /// Configure OAuth 2.1 Protected Resource Metadata for this transport.
283    ///
284    /// When set, adds a `GET` endpoint at the resource's path-aware RFC 9728
285    /// well-known location. This method only serves metadata; prefer
286    /// [`Self::into_oauth_router`] for a complete protected setup.
287    ///
288    /// # Example
289    ///
290    /// ```rust,ignore
291    /// use tower_mcp::oauth::ProtectedResourceMetadata;
292    /// use tower_mcp::transport::websocket::WebSocketTransport;
293    /// use tower_mcp::McpRouter;
294    ///
295    /// let metadata = ProtectedResourceMetadata::new("https://mcp.example.com")
296    ///     .authorization_server("https://auth.example.com")
297    ///     .scope("mcp:read");
298    ///
299    /// let router = McpRouter::new().server_info("my-server", "1.0.0");
300    /// let transport = WebSocketTransport::new(router).oauth(metadata);
301    /// ```
302    #[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    /// Build a fully protected OAuth WebSocket resource-server router.
309    ///
310    /// This validates and serves Protected Resource Metadata, authenticates
311    /// WebSocket upgrades, enforces the canonical resource audience, and
312    /// installs fail-closed per-operation scope checks.
313    #[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    /// Build a path-mounted, fully protected OAuth WebSocket router.
334    #[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    /// Apply a tower middleware layer to MCP request processing.
356    ///
357    /// The layer is applied to the [`McpRouter`] service within each session,
358    /// wrapping the `Service<RouterRequest>` pipeline. This allows middleware
359    /// like timeouts, rate limiting, or custom instrumentation to be applied
360    /// at the MCP request level.
361    ///
362    /// Middleware errors are automatically converted into JSON-RPC error
363    /// responses, so the transport's error handling remains unchanged.
364    ///
365    /// # Example
366    ///
367    /// ```rust,no_run
368    /// use std::time::Duration;
369    /// use tower::ServiceBuilder;
370    /// use tower::timeout::TimeoutLayer;
371    /// use tower_mcp::McpRouter;
372    /// use tower_mcp::transport::websocket::WebSocketTransport;
373    ///
374    /// let router = McpRouter::new().server_info("my-server", "1.0.0");
375    /// let transport = WebSocketTransport::new(router)
376    ///     .layer(
377    ///         ServiceBuilder::new()
378    ///             .layer(TimeoutLayer::new(Duration::from_secs(30)))
379    ///             .concurrency_limit(10)
380    ///             .into_inner(),
381    ///     );
382    /// ```
383    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    /// Build the axum router for this transport
403    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    /// Build an axum router mounted at a specific path
426    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    /// Serve the transport on the given address
451    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/// Add the OAuth Protected Resource Metadata well-known route if configured.
468#[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/// Parsed MCP WebSocket subprotocols from `Sec-WebSocket-Protocol` header.
496///
497/// Per SEP-1288, clients send `mcp.auth.{token}` and `mcp.version.{version}`
498/// as WebSocket subprotocols for authentication and version negotiation.
499#[derive(Debug, Default)]
500struct McpSubprotocols {
501    /// Authentication token extracted from `mcp.auth.{token}` subprotocol.
502    auth_token: Option<String>,
503    /// Protocol version extracted from `mcp.version.{version}` subprotocol.
504    protocol_version: Option<String>,
505    /// All matched subprotocol strings to echo back in the upgrade response.
506    selected: Vec<String>,
507}
508
509/// Parse MCP subprotocols from the `Sec-WebSocket-Protocol` header.
510///
511/// Returns the parsed subprotocols and the negotiated protocol version (if valid).
512fn 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
544/// Handle WebSocket upgrade.
545///
546/// Uses a raw `Request` extractor and performs the WebSocket upgrade manually
547/// so we can access HTTP request extensions (e.g., `TokenClaims` from OAuth
548/// middleware) and parse MCP subprotocols before upgrading.
549async 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    // Parse MCP subprotocols (mcp.auth.*, mcp.version.*) from Sec-WebSocket-Protocol
559    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    // Bridge TokenClaims from HTTP extensions to MCP extensions
565    #[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    // Store subprotocol auth token in extensions for downstream use
575    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    // Perform the WebSocket upgrade from request parts
585    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    // Echo back the matched subprotocols in the upgrade response
591    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/// Auth token extracted from the `mcp.auth.{token}` WebSocket subprotocol.
601///
602/// This is inserted into the MCP extensions map and can be accessed by
603/// middleware or tool handlers via `Extensions::get::<WebSocketAuthToken>()`.
604#[derive(Debug, Clone)]
605pub struct WebSocketAuthToken(pub String);
606
607/// Handle an individual WebSocket connection
608async fn handle_socket(
609    socket: WebSocket,
610    state: Arc<AppState>,
611    mcp_extensions: crate::router::Extensions,
612) {
613    // Use with_fresh_session() to ensure each session has its own state
614    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    // Cleanup session
649    state.sessions.remove(&session_id).await;
650    tracing::info!(session_id = %session_id, "WebSocket connection closed");
651}
652
653/// Handle WebSocket connection without sampling (simple mode)
654async 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    // Process incoming messages, also watching for cancellation (zombie prevention)
668    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                        // Notification, no response needed
715                    }
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                // MCP spec (SEP-1288) requires text frames only.
730                // Binary frames MUST result in close code 1003 (Unsupported Data).
731                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                // Ignore pongs
748            }
749            Message::Close(_) => {
750                tracing::info!(session_id = %session_id, "WebSocket close received");
751                break;
752            }
753        }
754    }
755}
756
757/// Handle WebSocket connection with sampling support (bidirectional mode)
758async 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    // Create channels for outgoing requests
767    let (request_tx, mut request_rx): (OutgoingRequestSender, OutgoingRequestReceiver) =
768        outgoing_request_channel(32);
769
770    // Create client requester for the router
771    let client_requester: ClientRequesterHandle = Arc::new(ChannelClientRequester::new(request_tx));
772
773    // Clone router and configure with client requester
774    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    // Track pending outgoing requests
783    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            // Handle incoming messages from client
794            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                        // MCP spec (SEP-1288) requires text frames only.
819                        // Binary frames MUST result in close code 1003 (Unsupported Data).
820                        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            // Handle outgoing requests to send to client
844            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            // Handle cancellation (zombie prevention)
856            _ = 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
871/// Handle an incoming WebSocket message (bidirectional mode)
872async 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    // Check if this is a response to one of our pending requests
901    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    // Check if it's a notification (no id field)
908    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(&notification)?;
911            router.handle_notification(mcp_notification);
912        }
913        return Ok(());
914    }
915
916    // Process as a request
917    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
942/// Handle a response to one of our pending requests
943async 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            // Send result to waiter (ignore if they've dropped the receiver)
990            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
1000/// Send an outgoing request to the client
1001async 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    // Build JSON-RPC request
1011    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    // Store pending request
1024    {
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    // Send the request
1035    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
1044/// Process a JSON-RPC message
1045async fn process_message(
1046    service: &mut JsonRpcService<McpBoxService>,
1047    router: &McpRouter,
1048    text: &str,
1049) -> Result<Option<crate::protocol::JsonRpcResponseMessage>> {
1050    // Check if it's a notification (no id field)
1051    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(&notification)?;
1064        router.handle_notification(mcp_notification);
1065        return Ok(None);
1066    }
1067
1068    // Parse and process as a request
1069    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        // Verify that .layer() compiles and produces a working transport
1119        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        // Only MCP subprotocols are selected
1222        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        // Should not be cancelled initially
1334        assert!(!*rx.borrow());
1335
1336        // After replace_connection, old receiver should see cancellation
1337        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        // First connection
1348        let _rx1 = session.cancel_receiver().await;
1349
1350        // Replace: old connection cancelled, new starts clean
1351        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        // Reconnect should cancel the first connection
1366        let result = store.reconnect(&session_id).await;
1367        assert!(result.is_some());
1368        let (_session2, rx2) = result.unwrap();
1369
1370        // Old receiver should see cancellation
1371        rx1.changed().await.unwrap();
1372        assert!(*rx1.borrow());
1373
1374        // New receiver should be clean
1375        assert!(!*rx2.borrow());
1376    }
1377}