Skip to main content

rama_ws/handshake/
client.rs

1//! WebSocket client types and utilities
2
3#![expect(
4    clippy::unreachable,
5    reason = "vendored from upstream `tungstenite-rs`: arms gated on caller-validated WebSocket protocol state that the type system can't enforce"
6)]
7
8use std::{
9    fmt,
10    future::Future,
11    ops::{Deref, DerefMut},
12    pin::Pin,
13    sync::Arc,
14    task::{Context, Poll},
15};
16
17use rama_core::Service;
18use rama_core::error::{BoxError, ErrorContext, ErrorExt};
19use rama_core::extensions::{Extensions, ExtensionsRef};
20use rama_core::futures::{Sink, SinkExt as _, Stream, StreamExt as _};
21use rama_core::rt::blocking::Io as BlockingIo;
22use rama_core::telemetry::tracing;
23use rama_http::conn::TargetHttpVersion;
24use rama_http::headers::sec_websocket_extensions::{Extension, PerMessageDeflateConfig};
25use rama_http::headers::sec_websocket_protocol::AcceptedWebSocketProtocol;
26use rama_http::headers::{
27    HeaderMapExt, HttpRequestBuilderExt as _, SecWebSocketExtensions, SecWebSocketKey,
28    SecWebSocketProtocol,
29};
30use rama_http::proto::h2::ext::Protocol;
31use rama_http::service::client::blocking::Client as BlockingHttpClient;
32use rama_http::service::client::ext::{IntoHeaderName, IntoHeaderValue};
33use rama_http::service::client::{HttpClientExt, IntoUrl, RequestBuilder};
34use rama_http::{Body, Method, Request, Response, StatusCode, Version, header, headers};
35use rama_http::{request, response};
36use rama_net::extensions::StreamTransformed;
37use rama_utils::str::NonEmptyStr;
38
39use crate::protocol::{CloseFrame, Message, ProtocolError, Role, WebSocket, WebSocketConfig};
40use crate::runtime::AsyncWebSocket;
41
42/// Builder that can be used by clients to initiate the WebSocket handshake.
43#[derive(Debug, Clone)]
44pub struct WebSocketRequestBuilder<B> {
45    inner: B,
46    protocols: Option<SecWebSocketProtocol>,
47    extensions: Option<SecWebSocketExtensions>,
48    key: Option<SecWebSocketKey>,
49}
50
51#[derive(Debug)]
52/// Request data to be used by an http client to initiate an http request.
53pub struct HandshakeRequest {
54    pub request: Request,
55    pub protocols: Option<SecWebSocketProtocol>,
56    pub extensions: Option<SecWebSocketExtensions>,
57    pub key: Option<SecWebSocketKey>,
58}
59
60struct PreparedHandshakeRequest {
61    request: Request,
62    protocols: Option<SecWebSocketProtocol>,
63    extensions: Option<SecWebSocketExtensions>,
64    config: Option<WebSocketConfig>,
65    key: Option<SecWebSocketKey>,
66}
67
68impl PreparedHandshakeRequest {
69    async fn send<S, Body>(
70        self,
71        service: &S,
72    ) -> Result<NegotiatedHandshakeRequest<Body>, HandshakeError>
73    where
74        S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
75    {
76        let uri = self.request.uri().clone();
77        let response = service.serve(self.request).await.map_err(|err| {
78            let err: BoxError = err.into();
79            HandshakeError::HttpRequestError(
80                err.context(uri)
81                    .context("send initial websocket handshake request (upgrade)"),
82            )
83        })?;
84
85        Ok(NegotiatedHandshakeRequest {
86            protocols: self.protocols,
87            extensions: self.extensions,
88            config: self.config,
89            key: self.key,
90            response,
91        })
92    }
93}
94
95/// [`WebSocketRequestBuilder`] inner wrapper type used for a builder,
96/// which includes a service, and thus is there to actually send the request as well and
97/// even follow up.
98pub struct WithService<'a, S, Body, Mode = websocket_builder_mode::Async> {
99    service: &'a S,
100    builder: RequestBuilder<'a, S, Response<Body>>,
101    config: Option<WebSocketConfig>,
102    is_h2: bool,
103    mode: Mode,
104}
105
106impl<S: fmt::Debug, Body, Mode: fmt::Debug> fmt::Debug for WithService<'_, S, Body, Mode> {
107    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
108        f.debug_struct("WithService")
109            .field("builder", &self.builder)
110            .field("config", &self.config)
111            .field("is_h2", &self.is_h2)
112            .field("mode", &self.mode)
113            .finish()
114    }
115}
116
117/// WebSocket request-builder execution modes.
118pub mod websocket_builder_mode {
119    use std::sync::Arc;
120
121    use rama_core::rt::blocking::Runtime;
122
123    /// Asynchronous terminal handshake operations.
124    #[derive(Debug)]
125    #[non_exhaustive]
126    pub struct Async;
127
128    /// Blocking terminal handshake operations.
129    #[derive(Debug, Clone)]
130    #[non_exhaustive]
131    pub struct Blocking<S> {
132        pub(crate) runtime: Runtime,
133        pub(crate) service: Arc<S>,
134    }
135}
136
137/// A WebSocket request builder whose terminal handshake operations block the
138/// calling thread.
139pub type BlockingWebSocketRequestBuilder<'a, S, Body> =
140    WebSocketRequestBuilder<WithService<'a, S, Body, websocket_builder_mode::Blocking<S>>>;
141
142fn new_ws_request_builder_from_uri<T>(uri: T, version: Version) -> request::Builder
143where
144    T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
145{
146    let builder = Request::builder()
147        .version(version)
148        .uri(uri)
149        .typed_header(headers::SecWebSocketVersion::V13);
150
151    match version {
152        version @ (Version::HTTP_10 | Version::HTTP_11) => builder
153            .method(Method::GET)
154            .version(version)
155            .typed_header(headers::Upgrade::websocket())
156            .typed_header(headers::Connection::upgrade()),
157        Version::HTTP_2 => builder.method(Method::CONNECT).version(Version::HTTP_2),
158        _ => unreachable!("bug"),
159    }
160}
161
162fn new_ws_request_builder_from_uri_with_service<'a, S, Body, T>(
163    service: &'a S,
164    uri: T,
165    version: Version,
166) -> RequestBuilder<'a, S, Response<Body>>
167where
168    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
169    T: IntoUrl,
170{
171    let builder = match version {
172        version @ (Version::HTTP_10 | Version::HTTP_11) => service
173            .get(uri)
174            .version(version)
175            .typed_header(headers::Upgrade::websocket())
176            .typed_header(headers::Connection::upgrade()),
177        Version::HTTP_2 => service.connect(uri).version(Version::HTTP_2),
178        _ => unreachable!("bug"),
179    };
180
181    builder.typed_header(headers::SecWebSocketVersion::V13)
182}
183
184fn new_ws_request_builder_from_request<'a, S, Body, RequestBody>(
185    service: &'a S,
186    mut request: Request<RequestBody>,
187) -> RequestBuilder<'a, S, Response<Body>>
188where
189    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
190    RequestBody: Into<rama_http::Body>,
191{
192    if !request
193        .headers()
194        .contains_key(header::SEC_WEBSOCKET_VERSION)
195    {
196        request
197            .headers_mut()
198            .typed_insert(headers::SecWebSocketVersion::V13);
199    }
200
201    match request.version() {
202        Version::HTTP_10 | Version::HTTP_11 => {
203            if request.headers().get(header::UPGRADE).is_none() {
204                request
205                    .headers_mut()
206                    .typed_insert(headers::Upgrade::websocket());
207            }
208            if request.headers().get(header::CONNECTION).is_none() {
209                request
210                    .headers_mut()
211                    .typed_insert(headers::Connection::upgrade());
212            }
213        }
214        // - for h2: nothing to do
215        // - else: this will error downstream due to invalid version
216        _ => (),
217    }
218    service.build_from_request(request)
219}
220
221#[derive(Debug)]
222/// Client error which can be triggered in case the response validation failed
223pub enum ResponseValidateError {
224    UnexpectedStatusCode(StatusCode),
225    UnexpectedHttpVersion(Version),
226    MissingUpgradeWebSocketHeader,
227    MissingConnectionUpgradeHeader,
228    SecWebSocketAcceptKeyMismatch,
229    ProtocolMismatch(Option<NonEmptyStr>),
230    ExtensionMismatch(Option<Extension>),
231}
232
233#[derive(Debug)]
234/// Client error which can be triggered in case the handshake phase failed.
235pub enum HandshakeError {
236    ValidationError(ResponseValidateError),
237    HttpRequestError(BoxError),
238    HttpUpgradeError(BoxError),
239}
240
241impl fmt::Display for ResponseValidateError {
242    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
243        match self {
244            Self::UnexpectedStatusCode(status_code) => {
245                write!(f, "unexpected HTTP status code: {status_code}")
246            }
247            Self::UnexpectedHttpVersion(version) => {
248                write!(f, "unexpected HTTP version: {version:?}")
249            }
250            Self::MissingUpgradeWebSocketHeader => {
251                write!(f, "missing upgrade WebSocket header")
252            }
253            Self::MissingConnectionUpgradeHeader => {
254                write!(f, "missing connection upgrade header")
255            }
256            Self::SecWebSocketAcceptKeyMismatch => {
257                write!(f, "key mismatch for sec-websocket-accept header")
258            }
259            Self::ProtocolMismatch(protocol) => {
260                write!(f, "protocol mismatch: {protocol:?}")
261            }
262            Self::ExtensionMismatch(extension) => {
263                write!(f, "extension mismatch: {extension:?}")
264            }
265        }
266    }
267}
268
269impl std::error::Error for ResponseValidateError {}
270
271impl fmt::Display for HandshakeError {
272    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
273        match self {
274            Self::ValidationError(error) => {
275                write!(f, "response validation failed: {error}")
276            }
277            Self::HttpRequestError(error) => {
278                write!(f, "http request error: {error}")
279            }
280            Self::HttpUpgradeError(error) => {
281                write!(f, "http upgrade error: {error}")
282            }
283        }
284    }
285}
286
287impl std::error::Error for HandshakeError {
288    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
289        match self {
290            Self::ValidationError(error) => Some(error as &dyn std::error::Error),
291            Self::HttpRequestError(error) | Self::HttpUpgradeError(error) => error.source(),
292        }
293    }
294}
295
296#[derive(Default, Debug)]
297pub struct AcceptedWebSocketData {
298    pub protocol: Option<AcceptedWebSocketProtocol>,
299    pub extension: Option<Extension>,
300}
301
302/// Validate the "accept" response from the http server
303/// with whom the client is trying to establish a WebSocket connection.
304pub fn validate_http_server_response<Body>(
305    response: &Response<Body>,
306    key: Option<headers::SecWebSocketKey>,
307    protocols: Option<SecWebSocketProtocol>,
308    extensions: Option<SecWebSocketExtensions>,
309) -> Result<AcceptedWebSocketData, ResponseValidateError> {
310    tracing::trace!(
311        http.version = ?response.version(),
312        http.response.status = ?response.status(),
313        ws.protocols = ?protocols,
314        ws.extensions = ?extensions,
315        "validate http server response"
316    );
317
318    match response.version() {
319        Version::HTTP_10 | Version::HTTP_11 => {
320            // If the status code received from the server is not 101, the
321            // client handles the response per HTTP [RFC2616] procedures. (RFC 6455)
322            let response_status = response.status();
323            if response_status != StatusCode::SWITCHING_PROTOCOLS {
324                return Err(ResponseValidateError::UnexpectedStatusCode(response_status));
325            }
326
327            // If the response lacks an |Upgrade| header field or the |Upgrade|
328            // header field contains a value that is not an ASCII case-
329            // insensitive match for the value "websocket", the client MUST
330            // _Fail the WebSocket Connection_. (RFC 6455)
331            if !response
332                .headers()
333                .typed_get::<headers::Upgrade>()
334                .map(|u| u.is_websocket())
335                .unwrap_or_default()
336            {
337                return Err(ResponseValidateError::MissingUpgradeWebSocketHeader);
338            }
339
340            // If the response lacks a |Connection| header field or the
341            // |Connection| header field doesn't contain a token that is an
342            // ASCII case-insensitive match for the value "Upgrade", the client
343            // MUST _Fail the WebSocket Connection_. (RFC 6455)
344            if !response
345                .headers()
346                .typed_get::<headers::Connection>()
347                .map(|c| c.contains_upgrade())
348                .unwrap_or_default()
349            {
350                return Err(ResponseValidateError::MissingConnectionUpgradeHeader);
351            }
352
353            // Sec-WebSocket-Key / Accept is only used in h1 responses.
354            //
355            // If the response lacks a |Sec-WebSocket-Accept| header field or
356            // the |Sec-WebSocket-Accept| contains a value other than the
357            // base64-encoded SHA-1 of ... the client MUST _Fail the WebSocket
358            // Connection_. (RFC 6455)
359            if let Some(key) = key {
360                let sec_websocket_accept_header = response
361                    .headers()
362                    .typed_get::<headers::SecWebSocketAccept>();
363                let expected_accept =
364                    headers::SecWebSocketAccept::try_from(key).map_err(|err| {
365                        tracing::debug!("failed to create WS accept header from key: {err}");
366                        ResponseValidateError::SecWebSocketAcceptKeyMismatch
367                    })?;
368                if sec_websocket_accept_header != Some(expected_accept) {
369                    tracing::trace!(
370                        "unexpected websocket accept key: {sec_websocket_accept_header:?}"
371                    );
372                    return Err(ResponseValidateError::SecWebSocketAcceptKeyMismatch);
373                }
374            }
375        }
376        Version::HTTP_2 => {
377            let response_status = response.status();
378            if !response.status().is_success() {
379                return Err(ResponseValidateError::UnexpectedStatusCode(response_status));
380            }
381        }
382        version => {
383            return Err(ResponseValidateError::UnexpectedHttpVersion(version));
384        }
385    }
386
387    // If the response includes a |Sec-WebSocket-Extensions| header
388    // field and this header field indicates the use of an extension
389    // that was not present in the client's handshake (the server has
390    // indicated an extension not requested by the client), the client
391    // MUST _Fail the WebSocket Connection_. (RFC 6455)
392    let mut accepted_extension = None;
393    match (
394        response
395            .headers()
396            .typed_get::<SecWebSocketExtensions>()
397            .map(|ext| ext.0.head),
398        extensions,
399    ) {
400        (None, Some(allowed_extensions)) => {
401            tracing::trace!(
402                ws.extensions = ?allowed_extensions,
403                "server selected no WS extensions despite client supporting some (valid, move on without)",
404            );
405        }
406        (Some(Extension::PerMessageDeflate(server_cfg)), Some(client_extensions)) => {
407            accepted_extension = client_extensions
408                .0.iter()
409                .find_map(|client_ext| {
410                    if let Extension::PerMessageDeflate(client_cfg) = client_ext {
411                        return Some(Ok(Extension::PerMessageDeflate(PerMessageDeflateConfig {
412                            client_max_window_bits: match (
413                                server_cfg.client_max_window_bits,
414                                client_cfg.client_max_window_bits,
415                            ) {
416                                (None, None | Some(_)) => None,
417                                (Some(srv), maybe_offered) => {
418                                    if !(8..=15).contains(&srv) || maybe_offered.map(|offered| offered != 0 && srv > offered).unwrap_or_default() {
419                                        tracing::debug!("server offered invalid client_max_window_bits (pmd)... ext mismatch!");
420                                        return Some(Err(
421                                            ResponseValidateError::ExtensionMismatch(Some(
422                                                Extension::PerMessageDeflate(server_cfg.clone()),
423                                            )),
424                                        ));
425                                    }
426                                    Some(srv)
427                                }
428                            },
429                            server_max_window_bits: match (
430                                server_cfg.server_max_window_bits,
431                                client_cfg.server_max_window_bits,
432                            ) {
433                                (None, None | Some(_)) => None,
434                                (Some(their_bits), maybe_our_bits) => {
435                                    if !(8..=15).contains(&their_bits)
436                                        || maybe_our_bits
437                                            .map(|our_bits| our_bits != 0 && their_bits > our_bits)
438                                            .unwrap_or_default()
439                                    {
440                                        tracing::debug!("server offered invalid server_max_window_bits (pmd)... ext mismatch!");
441                                        return Some(Err(
442                                            ResponseValidateError::ExtensionMismatch(Some(
443                                                Extension::PerMessageDeflate(server_cfg.clone()),
444                                            )),
445                                        ));
446                                    }
447                                    Some(their_bits)
448                                }
449                            },
450                            server_no_context_takeover: server_cfg.server_no_context_takeover,
451                            client_no_context_takeover: client_cfg.client_no_context_takeover,
452                            identifier: server_cfg.identifier.clone(),
453                        })));
454                    }
455                    None
456                })
457                .transpose()?;
458        }
459        (Some(server_ext), _) => {
460            tracing::debug!("server offered ext, but client (we) not!");
461            return Err(ResponseValidateError::ExtensionMismatch(Some(server_ext)));
462        }
463        (None, None) => (),
464    }
465
466    // If the response includes a |Sec-WebSocket-Protocol| header field
467    // and this header field indicates the use of a subprotocol that was
468    // not present in the client's handshake (the server has indicated a
469    // subprotocol not requested by the client), the client MUST _Fail
470    // the WebSocket Connection_. (RFC 6455)
471    let mut accepted_protocol = None;
472    match (
473        response
474            .headers()
475            .typed_get::<SecWebSocketProtocol>()
476            .map(|h| h.accept_first_protocol()),
477        protocols,
478    ) {
479        (None, None) => (),
480        (None, Some(allowed_protocols)) => {
481            // RFC 6455 only mandates failure when the server selects a protocol
482            // not in the client's offer — a server may legitimately decline to
483            // select any subprotocol even when the client proposed one.
484            tracing::trace!(
485                ws.protocols = ?allowed_protocols,
486                "server selected no WS subprotocol despite client proposing some (valid, proceed without)",
487            );
488        }
489        (Some(header), None) => {
490            return Err(ResponseValidateError::ProtocolMismatch(Some(header.0)));
491        }
492        (Some(protocol_header), Some(sub_protocols)) => {
493            match sub_protocols.contains(&protocol_header.0) {
494                Some(protocol) => accepted_protocol = Some(protocol),
495                None => {
496                    return Err(ResponseValidateError::ProtocolMismatch(Some(
497                        protocol_header.0,
498                    )));
499                }
500            };
501        }
502    }
503
504    Ok(AcceptedWebSocketData {
505        protocol: accepted_protocol,
506        extension: accepted_extension,
507    })
508}
509
510impl WebSocketRequestBuilder<request::Builder> {
511    /// Create a new `http/1.1` WebSocket [`Request`] builder.
512    pub fn new<T>(uri: T) -> Self
513    where
514        T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
515    {
516        Self::new_with_version(uri, Version::HTTP_11)
517    }
518
519    /// Create a new `h2` WebSocket [`Request`] builder.
520    pub fn new_h2<T>(uri: T) -> Self
521    where
522        T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
523    {
524        Self::new_with_version(uri, Version::HTTP_2)
525    }
526
527    fn new_with_version<T>(uri: T, version: Version) -> Self
528    where
529        T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
530    {
531        Self {
532            inner: new_ws_request_builder_from_uri(uri, version),
533            protocols: Default::default(),
534            extensions: Default::default(),
535            key: Default::default(),
536        }
537    }
538
539    /// Set a custom http header
540    #[must_use]
541    pub fn with_header<K, V>(self, name: K, value: V) -> Self
542    where
543        K: TryInto<rama_http::HeaderName, Error: Into<rama_http::HttpError>>,
544        V: TryInto<rama_http::HeaderValue, Error: Into<rama_http::HttpError>>,
545    {
546        Self {
547            inner: self.inner.header(name, value),
548            protocols: self.protocols,
549            extensions: self.extensions,
550            key: self.key,
551        }
552    }
553
554    /// Set a custom typed http header
555    #[must_use]
556    pub fn with_typed_header<H>(self, header: H) -> Self
557    where
558        H: headers::HeaderEncode,
559    {
560        Self {
561            inner: self.inner.typed_header(header),
562            protocols: self.protocols,
563            extensions: self.extensions,
564            key: self.key,
565        }
566    }
567
568    /// Build the handshake data
569    /// to be used to initiate the WebSocket handshake using an http client.
570    pub fn build_handshake(self) -> Result<HandshakeRequest, BoxError> {
571        let builder = match self.protocols.as_ref() {
572            Some(protocols) => self.inner.typed_header(protocols),
573            None => self.inner,
574        };
575
576        let builder = match self.extensions.as_ref() {
577            Some(extensions) => builder.typed_header(extensions),
578            None => builder,
579        };
580
581        let mut request = builder
582            .body(Body::empty())
583            .context("request failed to build (invalid custom header?)")?;
584
585        let mut key = None;
586        if request.version() != Version::HTTP_2 {
587            let k = self.key.unwrap_or_else(headers::SecWebSocketKey::random);
588            request.headers_mut().typed_insert(&k);
589            key = Some(k);
590        }
591
592        // only required for h2, but we might upgrade from h1 to h2 based on layers such as tls
593        request
594            .extensions()
595            .insert(Protocol::from_static("websocket"));
596
597        Ok(HandshakeRequest {
598            request,
599            protocols: self.protocols,
600            extensions: self.extensions,
601            key,
602        })
603    }
604}
605
606impl<'a, S, Body> WebSocketRequestBuilder<WithService<'a, S, Body, websocket_builder_mode::Async>>
607where
608    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
609{
610    /// Create a new `http/1.1` WebSocket [`Request`] builder.
611    pub fn new_with_service<T>(service: &'a S, uri: T) -> Self
612    where
613        T: IntoUrl,
614    {
615        Self::new_with_service_and_version_and_mode(
616            service,
617            Version::HTTP_11,
618            uri,
619            websocket_builder_mode::Async,
620        )
621    }
622
623    /// Create a new `h2` WebSocket [`Request`] builder.
624    pub fn new_h2_with_service<T>(service: &'a S, uri: T) -> Self
625    where
626        T: IntoUrl,
627    {
628        Self::new_with_service_and_version_and_mode(
629            service,
630            Version::HTTP_2,
631            uri,
632            websocket_builder_mode::Async,
633        )
634    }
635
636    /// Create a new WebSocket [`Request`] builder for the given [`Request`]
637    pub fn new_with_service_and_request<RequestBody>(
638        service: &'a S,
639        request: Request<RequestBody>,
640    ) -> Self
641    where
642        RequestBody: Into<rama_http::Body>,
643    {
644        Self::new_with_service_request_and_mode(service, request, websocket_builder_mode::Async)
645    }
646}
647
648impl<'a, S, Body, Mode> WebSocketRequestBuilder<WithService<'a, S, Body, Mode>>
649where
650    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
651{
652    fn new_with_service_and_version_and_mode<T>(
653        service: &'a S,
654        version: Version,
655        uri: T,
656        mode: Mode,
657    ) -> Self
658    where
659        T: IntoUrl,
660    {
661        Self {
662            inner: WithService {
663                service,
664                builder: new_ws_request_builder_from_uri_with_service(service, uri, version),
665                config: Default::default(),
666                is_h2: version == Version::HTTP_2,
667                mode,
668            },
669            protocols: Default::default(),
670            extensions: Default::default(),
671            key: Default::default(),
672        }
673    }
674
675    fn new_with_service_request_and_mode<RequestBody>(
676        service: &'a S,
677        request: Request<RequestBody>,
678        mode: Mode,
679    ) -> Self
680    where
681        RequestBody: Into<rama_http::Body>,
682    {
683        let key = request.headers().typed_get();
684        let is_h2 = request.version() == Version::HTTP_2;
685        let protocols = request.headers().typed_get();
686        let extensions = request.headers().typed_get();
687
688        Self {
689            inner: WithService {
690                service,
691                builder: new_ws_request_builder_from_request(service, request),
692                config: Default::default(),
693                is_h2,
694                mode,
695            },
696            protocols,
697            extensions,
698            key,
699        }
700    }
701
702    /// Set a custom http header
703    #[must_use]
704    pub fn with_header<K, V>(self, name: K, value: V) -> Self
705    where
706        K: IntoHeaderName,
707        V: IntoHeaderValue,
708    {
709        Self {
710            inner: WithService {
711                builder: self.inner.builder.header(name, value),
712                ..self.inner
713            },
714            protocols: self.protocols,
715            extensions: self.extensions,
716            key: self.key,
717        }
718    }
719
720    /// Overwrite a custom http header
721    #[must_use]
722    pub fn with_header_overwrite<K, V>(self, name: K, value: V) -> Self
723    where
724        K: IntoHeaderName,
725        V: IntoHeaderValue,
726    {
727        Self {
728            inner: WithService {
729                builder: self.inner.builder.overwrite_header(name, value),
730                ..self.inner
731            },
732            protocols: self.protocols,
733            extensions: self.extensions,
734            key: self.key,
735        }
736    }
737
738    /// Set a custom typed http header
739    #[must_use]
740    pub fn with_typed_header<H>(self, header: H) -> Self
741    where
742        H: headers::HeaderEncode,
743    {
744        Self {
745            inner: WithService {
746                builder: self.inner.builder.typed_header(header),
747                ..self.inner
748            },
749            protocols: self.protocols,
750            extensions: self.extensions,
751            key: self.key,
752        }
753    }
754
755    /// Overwrite a custom typed http header
756    #[must_use]
757    pub fn with_typed_header_overwrite<H>(self, header: H) -> Self
758    where
759        H: headers::HeaderEncode,
760    {
761        Self {
762            inner: WithService {
763                builder: self.inner.builder.overwrite_typed_header(header),
764                ..self.inner
765            },
766            protocols: self.protocols,
767            extensions: self.extensions,
768            key: self.key,
769        }
770    }
771
772    #[cfg(feature = "compression")]
773    rama_utils::macros::generate_set_and_with! {
774        /// Set/add deflate ext and also apply it to the [`WebSocketConfig`],
775        /// using the default [`crate::protocol::PerMessageDeflateConfig`].
776        #[must_use]
777        #[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
778        pub fn per_message_deflate(mut self) -> Self {
779            self.extensions = match self.extensions.take() {
780                Some(ext) => {
781                    Some(ext.with_extra_extension(Extension::PerMessageDeflate(Default::default())))
782                },
783                None => Some(SecWebSocketExtensions::per_message_deflate()),
784            };
785            self.inner.config = Some(self.inner.config.take().unwrap_or_default().with_per_message_deflate_default());
786            self
787        }
788    }
789
790    #[cfg(feature = "compression")]
791    rama_utils::macros::generate_set_and_with! {
792        /// Set/add deflate ext and also apply it to the [`WebSocketConfig`],
793        /// using the default [`crate::protocol::PerMessageDeflateConfig`].
794        ///
795        /// Overwrites existing extensions if already existed.
796        #[must_use]
797        #[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
798        pub fn per_message_deflate_overwrite_extensions(mut self) -> Self {
799            self.extensions = Some(SecWebSocketExtensions::per_message_deflate());
800            self.inner.config = Some(self.inner.config.take().unwrap_or_default().with_per_message_deflate_default());
801            self
802        }
803    }
804
805    #[cfg(feature = "compression")]
806    rama_utils::macros::generate_set_and_with! {
807        /// Set/add deflate ext and also apply it to the [`WebSocketConfig`],
808        /// using the default [`crate::protocol::PerMessageDeflateConfig`].
809        #[must_use]
810        #[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
811        pub fn per_message_deflate_with_config(mut self, config: impl Into<crate::protocol::PerMessageDeflateConfig>) -> Self {
812            let config = config.into();
813            self.extensions = match self.extensions.take() {
814                Some(ext) => {
815                    Some(ext.with_extra_extension(Extension::PerMessageDeflate((&config).into())))
816                }
817                None => Some(SecWebSocketExtensions::per_message_deflate_with_config((&config).into())),
818            };
819            self.inner.config = Some(
820                self.inner
821                    .config
822                    .take()
823                    .unwrap_or_default()
824                    .with_per_message_deflate(config),
825            );
826            self
827        }
828    }
829
830    #[cfg(feature = "compression")]
831    rama_utils::macros::generate_set_and_with! {
832        /// Set/add deflate ext and also apply it to the [`WebSocketConfig`],
833        /// using the default [`crate::protocol::PerMessageDeflateConfig`].
834        ///
835        /// Overwrites existing extensions if already existed.
836        #[must_use]
837        #[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
838        pub fn per_message_deflate_with_config_overwrite_extensions(mut self, config: impl Into<crate::protocol::PerMessageDeflateConfig>) -> Self {
839            let config = config.into();
840            self.extensions = Some(SecWebSocketExtensions::per_message_deflate_with_config((&config).into()));
841            self.inner.config = Some(
842                self.inner
843                    .config
844                    .take()
845                    .unwrap_or_default()
846                    .with_per_message_deflate(config),
847            );
848            self
849        }
850    }
851
852    rama_utils::macros::generate_set_and_with! {
853        /// Set the [`WebSocketConfig`], overwriting the previous config if already set.
854        pub fn config(mut self, cfg: Option<WebSocketConfig>) -> Self {
855            self.inner.config = cfg;
856            self
857        }
858    }
859
860    fn prepare_handshake_inner(
861        self,
862        extensions: &Extensions,
863    ) -> Result<PreparedHandshakeRequest, HandshakeError> {
864        extensions.insert(StreamTransformed {
865            by: "rama-ws::WebSocketClient",
866        });
867
868        let builder = match self.protocols.as_ref() {
869            Some(protocols) => self.inner.builder.overwrite_typed_header(protocols),
870            None => self.inner.builder,
871        };
872
873        let builder = match self.extensions.as_ref() {
874            Some(extensions) => builder.typed_header(extensions),
875            None => builder,
876        };
877
878        let mut key = None;
879        let builder = if !self.inner.is_h2 {
880            extensions.insert(TargetHttpVersion(Version::HTTP_11));
881
882            let k = self.key.unwrap_or_else(headers::SecWebSocketKey::random);
883            let builder = builder.overwrite_typed_header(&k);
884            key = Some(k);
885            builder
886        } else {
887            extensions.insert(TargetHttpVersion(Version::HTTP_2));
888
889            builder
890        };
891
892        // only required in h1, but because of layers such as tls we might anyway turn from h1 into h2
893        let builder = builder.extension(Protocol::from_static("websocket"));
894
895        if let Some(ext) = builder.extensions() {
896            ext.extend(extensions);
897        }
898
899        let request = builder
900            .build()
901            .context("build initial websocket handshake request (upgrade)")
902            .map_err(HandshakeError::HttpRequestError)?;
903
904        Ok(PreparedHandshakeRequest {
905            request,
906            protocols: self.protocols,
907            extensions: self.extensions,
908            config: self.inner.config,
909            key,
910        })
911    }
912
913    async fn initiate_handshake_inner(
914        self,
915        extensions: Extensions,
916    ) -> Result<NegotiatedHandshakeRequest<Body>, HandshakeError> {
917        let service = self.inner.service;
918        let prepared = self.prepare_handshake_inner(&extensions)?;
919        prepared.send(service).await
920    }
921}
922
923impl<'a, S, Body> WebSocketRequestBuilder<WithService<'a, S, Body, websocket_builder_mode::Async>>
924where
925    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
926{
927    /// Initiate the handshake by preparing the http request, sending it
928    /// and receiving the http response.
929    ///
930    /// This consumes this [`WebSocketRequestBuilder`]. Fulfill
931    /// the handshake by calling [`NegotiatedHandshakeRequest::complete`].
932    ///
933    /// In most cases you have however no need for this intermediate result,
934    /// and are better of calling [`Self::handshake`] directly. Only in cases
935    /// such as MITM proxies or edge-case purposes you might require access
936    /// to [`NegotiatedHandshakeRequest`].
937    pub async fn initiate_handshake(
938        self,
939        extensions: Extensions,
940    ) -> Result<NegotiatedHandshakeRequest<Body>, HandshakeError> {
941        self.initiate_handshake_inner(extensions).await
942    }
943
944    /// Establish a [`ClientWebSocket`], consuming this [`WebSocketRequestBuilder`],
945    /// by doing the http-handshake, including validation and returning the socket if all is good.
946    pub async fn handshake(self, extensions: Extensions) -> Result<ClientWebSocket, HandshakeError>
947    where
948        Body: Send + 'static,
949    {
950        let handshake = self.initiate_handshake(extensions).await?;
951        handshake.complete().await
952    }
953}
954
955impl<'a, S, Body>
956    WebSocketRequestBuilder<WithService<'a, S, Body, websocket_builder_mode::Blocking<S>>>
957where
958    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
959    Body: Send + 'static,
960{
961    fn new_blocking_with_service<T>(client: &'a BlockingHttpClient<S>, uri: T) -> Self
962    where
963        T: IntoUrl,
964    {
965        Self::new_with_service_and_version_and_mode(
966            client.get_ref(),
967            Version::HTTP_11,
968            uri,
969            websocket_builder_mode::Blocking {
970                runtime: client.runtime().clone(),
971                service: client.clone_service(),
972            },
973        )
974    }
975
976    fn new_blocking_h2_with_service<T>(client: &'a BlockingHttpClient<S>, uri: T) -> Self
977    where
978        T: IntoUrl,
979    {
980        Self::new_with_service_and_version_and_mode(
981            client.get_ref(),
982            Version::HTTP_2,
983            uri,
984            websocket_builder_mode::Blocking {
985                runtime: client.runtime().clone(),
986                service: client.clone_service(),
987            },
988        )
989    }
990
991    fn new_blocking_with_service_and_request<RequestBody>(
992        client: &'a BlockingHttpClient<S>,
993        request: Request<RequestBody>,
994    ) -> Self
995    where
996        RequestBody: Into<rama_http::Body>,
997    {
998        Self::new_with_service_request_and_mode(
999            client.get_ref(),
1000            request,
1001            websocket_builder_mode::Blocking {
1002                runtime: client.runtime().clone(),
1003                service: client.clone_service(),
1004            },
1005        )
1006    }
1007
1008    /// Establish a blocking [`BlockingClientWebSocket`] using empty request
1009    /// extensions.
1010    pub fn try_handshake(self) -> Result<BlockingClientWebSocket, HandshakeError> {
1011        self.try_handshake_with_extensions(Extensions::new())
1012    }
1013
1014    /// Establish a blocking [`BlockingClientWebSocket`] using the supplied
1015    /// request extensions.
1016    #[expect(
1017        clippy::needless_pass_by_value,
1018        reason = "matches the async handshake API and transfers the extension set"
1019    )]
1020    pub fn try_handshake_with_extensions(
1021        self,
1022        extensions: Extensions,
1023    ) -> Result<BlockingClientWebSocket, HandshakeError> {
1024        let runtime = self.inner.mode.runtime.clone();
1025        let service = Arc::clone(&self.inner.mode.service);
1026        let prepared = self.prepare_handshake_inner(&extensions)?;
1027        let completed = runtime.block_on_task(async move {
1028            let handshake = prepared.send(service.as_ref()).await?;
1029            handshake.complete_upgrade().await
1030        })?;
1031        let socket = WebSocket::from_raw_socket(
1032            runtime.io(completed.stream),
1033            Role::Client,
1034            completed.config,
1035        );
1036
1037        Ok(BlockingClientWebSocket {
1038            socket,
1039            response: completed.response,
1040            accepted_protocol: completed.accepted_protocol,
1041        })
1042    }
1043}
1044
1045impl<B> WebSocketRequestBuilder<B> {
1046    rama_utils::macros::generate_set_and_with! {
1047        /// Define the WebSocket protocols to be used.
1048        pub fn protocols(mut self, protocols: Option<SecWebSocketProtocol>) -> Self {
1049            self.protocols = protocols;
1050            self
1051        }
1052    }
1053
1054    rama_utils::macros::generate_set_and_with! {
1055        /// Set the WebSocket key (a random one will be generated if not defined).
1056        ///
1057        /// Only touch this property if you have a good reason to do so.
1058        pub fn key(mut self, key: Option<headers::SecWebSocketKey>) -> Self {
1059            self.key = key;
1060            self
1061        }
1062    }
1063}
1064
1065/// Utility which can be used my Mitm proxies to
1066/// update the base config of a client websocket config.
1067pub fn apply_response_data_to_base_websocket_config<Body>(
1068    base_cfg: Option<WebSocketConfig>,
1069    res: &mut Response<Body>,
1070) -> Option<WebSocketConfig> {
1071    let accepted_pmd_cfg = res
1072        .headers()
1073        .typed_get::<SecWebSocketExtensions>()
1074        .map(|ext| ext.0.head)
1075        .and_then(|ext| {
1076            if let Extension::PerMessageDeflate(cfg) = ext {
1077                Some(cfg)
1078            } else {
1079                None
1080            }
1081        });
1082
1083    if let Some(accepted_protocol) = res
1084        .headers()
1085        .typed_get::<SecWebSocketProtocol>()
1086        .map(|h| h.accept_first_protocol())
1087    {
1088        res.extensions().insert(accepted_protocol);
1089    }
1090
1091    #[cfg(feature = "compression")]
1092    {
1093        if let Some(pmd_cfg) = accepted_pmd_cfg {
1094            let mut ws_cfg = base_cfg.unwrap_or_default();
1095            ws_cfg.per_message_deflate = Some(pmd_cfg.into());
1096            Some(ws_cfg)
1097        } else if let Some(mut ws_cfg) = base_cfg {
1098            ws_cfg.per_message_deflate = None;
1099            Some(ws_cfg)
1100        } else {
1101            base_cfg
1102        }
1103    }
1104
1105    #[cfg(not(feature = "compression"))]
1106    {
1107        if accepted_pmd_cfg.is_some() {
1108            tracing::error!(
1109                "per-message-deflate is used but compression feature is disabled. Enable it if you wish to use this extension."
1110            );
1111        }
1112
1113        base_cfg
1114    }
1115}
1116
1117/// Intermediate websocket handshake created by
1118/// [`WebSocketRequestBuilder::initiate_handshake`].
1119///
1120/// Useful in case you require access to some of the data
1121/// prior to validation and WS upgrading.
1122pub struct NegotiatedHandshakeRequest<Body> {
1123    pub protocols: Option<SecWebSocketProtocol>,
1124    pub extensions: Option<SecWebSocketExtensions>,
1125    pub config: Option<WebSocketConfig>,
1126    pub key: Option<SecWebSocketKey>,
1127    pub response: Response<Body>,
1128}
1129
1130struct CompletedClientHandshake {
1131    stream: rama_http::io::upgrade::Upgraded,
1132    response: response::Parts,
1133    accepted_protocol: Option<AcceptedWebSocketProtocol>,
1134    config: Option<WebSocketConfig>,
1135}
1136
1137impl<Body> NegotiatedHandshakeRequest<Body> {
1138    /// Fulfill the websocket handshake and return the upgraded [`ClientWebSocket`].
1139    pub async fn complete(self) -> Result<ClientWebSocket, HandshakeError>
1140    where
1141        Body: Send + 'static,
1142    {
1143        let completed = self.complete_upgrade().await?;
1144        let socket =
1145            AsyncWebSocket::from_raw_socket(completed.stream, Role::Client, completed.config).await;
1146
1147        Ok(ClientWebSocket {
1148            socket,
1149            response: completed.response,
1150            accepted_protocol: completed.accepted_protocol,
1151        })
1152    }
1153
1154    async fn complete_upgrade(self) -> Result<CompletedClientHandshake, HandshakeError>
1155    where
1156        Body: Send + 'static,
1157    {
1158        let accepted_data = validate_http_server_response(
1159            &self.response,
1160            self.key,
1161            self.protocols,
1162            self.extensions,
1163        )
1164        .map_err(HandshakeError::ValidationError)?;
1165
1166        tracing::trace!(
1167            websocket.protocol = ?accepted_data.protocol,
1168            websocket.extension = ?accepted_data.extension,
1169            "websocket handshake http response is valid",
1170        );
1171
1172        #[cfg(feature = "compression")]
1173        let maybe_ws_cfg = {
1174            let mut ws_cfg = self.config.unwrap_or_default();
1175
1176            if let Some(Extension::PerMessageDeflate(pmd_cfg)) = accepted_data.extension {
1177                tracing::trace!(
1178                    "apply accepted per-message-deflate cfg into WS client config: {pmd_cfg:?}"
1179                );
1180                ws_cfg.per_message_deflate = Some(pmd_cfg.into());
1181            } else {
1182                ws_cfg.per_message_deflate = None;
1183            }
1184
1185            Some(ws_cfg)
1186        };
1187
1188        #[cfg(not(feature = "compression"))]
1189        let maybe_ws_cfg = {
1190            if let Some(Extension::PerMessageDeflate(pmd_cfg)) = accepted_data.extension {
1191                tracing::error!(
1192                    "per-message-deflate is used but compression feature is disabled. Enable it if you wish to use this extension."
1193                );
1194                return Err(HandshakeError::ValidationError(
1195                    ResponseValidateError::ExtensionMismatch(Some(Extension::PerMessageDeflate(
1196                        pmd_cfg,
1197                    ))),
1198                ));
1199            }
1200            self.config
1201        };
1202
1203        let on_upgrade = rama_http::io::upgrade::handle_upgrade(&self.response);
1204        let (parts, body) = self.response.into_parts();
1205        let stream = on_upgrade
1206            .await
1207            .context("upgrade http connection into a raw web socket")
1208            .map_err(HandshakeError::HttpUpgradeError)?
1209            .with_guard(body);
1210        Ok(CompletedClientHandshake {
1211            stream,
1212            response: parts,
1213            accepted_protocol: accepted_data.protocol,
1214            config: maybe_ws_cfg,
1215        })
1216    }
1217}
1218
1219#[derive(Debug)]
1220/// [`ClientWebSocket`], used as input-output stream.
1221///
1222/// Utility type created via [`WebSocketRequestBuilder::handshake`].
1223pub struct ClientWebSocket<S = AsyncWebSocket> {
1224    /// Established WebSocket message transport.
1225    pub socket: S,
1226    /// Original HTTP handshake response metadata.
1227    pub response: response::Parts,
1228    /// Subprotocol accepted during the HTTP handshake, when any.
1229    pub accepted_protocol: Option<AcceptedWebSocketProtocol>,
1230}
1231
1232impl<S> Deref for ClientWebSocket<S> {
1233    type Target = S;
1234
1235    fn deref(&self) -> &Self::Target {
1236        &self.socket
1237    }
1238}
1239
1240impl<S> DerefMut for ClientWebSocket<S> {
1241    fn deref_mut(&mut self) -> &mut Self::Target {
1242        &mut self.socket
1243    }
1244}
1245
1246impl<S> Stream for ClientWebSocket<S>
1247where
1248    S: Stream<Item = Result<Message, ProtocolError>> + Unpin,
1249{
1250    type Item = Result<Message, ProtocolError>;
1251
1252    fn poll_next(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1253        Stream::poll_next(Pin::new(&mut self.get_mut().socket), ctx)
1254    }
1255}
1256
1257impl<S> Sink<Message> for ClientWebSocket<S>
1258where
1259    S: Sink<Message, Error = ProtocolError> + Unpin,
1260{
1261    type Error = ProtocolError;
1262
1263    fn poll_ready(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1264        Sink::poll_ready(Pin::new(&mut self.get_mut().socket), ctx)
1265    }
1266
1267    fn start_send(self: Pin<&mut Self>, message: Message) -> Result<(), Self::Error> {
1268        Sink::start_send(Pin::new(&mut self.get_mut().socket), message)
1269    }
1270
1271    fn poll_flush(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1272        Sink::poll_flush(Pin::new(&mut self.get_mut().socket), ctx)
1273    }
1274
1275    fn poll_close(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1276        Sink::poll_close(Pin::new(&mut self.get_mut().socket), ctx)
1277    }
1278}
1279
1280impl<S> ExtensionsRef for ClientWebSocket<S>
1281where
1282    S: ExtensionsRef,
1283{
1284    fn extensions(&self) -> &Extensions {
1285        self.socket.extensions()
1286    }
1287}
1288
1289impl<S> ClientWebSocket<S> {
1290    /// Transform the message transport while preserving handshake metadata.
1291    #[must_use]
1292    pub fn map_socket<T>(self, map: impl FnOnce(S) -> T) -> ClientWebSocket<T> {
1293        ClientWebSocket {
1294            socket: map(self.socket),
1295            response: self.response,
1296            accepted_protocol: self.accepted_protocol,
1297        }
1298    }
1299
1300    /// Write and flush one message.
1301    pub fn send_message(
1302        &mut self,
1303        message: Message,
1304    ) -> impl Future<Output = Result<(), ProtocolError>> + Send + '_
1305    where
1306        S: Sink<Message, Error = ProtocolError> + Send + Unpin,
1307    {
1308        self.socket.send(message)
1309    }
1310
1311    /// Receive one complete message.
1312    pub async fn recv_message(&mut self) -> Result<Message, ProtocolError>
1313    where
1314        S: Stream<Item = Result<Message, ProtocolError>> + Unpin,
1315    {
1316        self.socket.next().await.ok_or_else(|| {
1317            ProtocolError::Io(std::io::Error::new(
1318                std::io::ErrorKind::ConnectionAborted,
1319                "Connection closed: no messages to receive",
1320            ))
1321        })?
1322    }
1323
1324    /// Close the WebSocket.
1325    pub async fn close(&mut self, message: Option<CloseFrame>) -> Result<(), ProtocolError>
1326    where
1327        S: Sink<Message, Error = ProtocolError> + Send + Unpin,
1328    {
1329        self.socket.send(Message::Close(message)).await
1330    }
1331
1332    /// View the original response data, from which this client web socket was created.
1333    pub fn response(&self) -> &response::Parts {
1334        &self.response
1335    }
1336
1337    /// Return the accepted protocol (during the http handshake) of the [`ClientWebSocket`], if any.
1338    pub fn accepted_protocol(&self) -> Option<&str> {
1339        self.accepted_protocol.as_ref().map(|p| p.0.as_ref())
1340    }
1341
1342    /// Consume `self` and return its message transport.
1343    pub fn into_inner(self) -> S {
1344        self.socket
1345    }
1346}
1347
1348/// A synchronous WebSocket over an upgraded HTTP transport driven by a Rama
1349/// blocking runtime.
1350pub type BlockingWebSocket = WebSocket<BlockingIo<rama_http::io::upgrade::Upgraded>>;
1351
1352/// A connected blocking client WebSocket and its HTTP handshake metadata.
1353#[derive(Debug)]
1354pub struct BlockingClientWebSocket {
1355    /// Established blocking WebSocket transport.
1356    pub socket: BlockingWebSocket,
1357    /// Original HTTP handshake response metadata.
1358    pub response: response::Parts,
1359    /// Subprotocol accepted during the HTTP handshake, when any.
1360    pub accepted_protocol: Option<AcceptedWebSocketProtocol>,
1361}
1362
1363impl Deref for BlockingClientWebSocket {
1364    type Target = BlockingWebSocket;
1365
1366    fn deref(&self) -> &Self::Target {
1367        &self.socket
1368    }
1369}
1370
1371impl DerefMut for BlockingClientWebSocket {
1372    fn deref_mut(&mut self) -> &mut Self::Target {
1373        &mut self.socket
1374    }
1375}
1376
1377impl BlockingClientWebSocket {
1378    /// View the original response data from which this WebSocket was created.
1379    pub fn response(&self) -> &response::Parts {
1380        &self.response
1381    }
1382
1383    /// Return the subprotocol accepted during the HTTP handshake, if any.
1384    pub fn accepted_protocol(&self) -> Option<&str> {
1385        self.accepted_protocol.as_ref().map(|p| p.0.as_ref())
1386    }
1387
1388    /// Write and immediately flush a message.
1389    pub fn send_message(&mut self, message: Message) -> Result<(), ProtocolError> {
1390        self.socket.send(message)
1391    }
1392
1393    /// Read the next message.
1394    pub fn recv_message(&mut self) -> Result<Message, ProtocolError> {
1395        self.socket.read()
1396    }
1397
1398    /// Consume this wrapper and return the blocking WebSocket.
1399    pub fn into_inner(self) -> BlockingWebSocket {
1400        self.socket
1401    }
1402}
1403
1404/// Extends an Http Client with high level features WebSocket features.
1405pub trait HttpClientWebSocketExt<Body>:
1406    private::HttpClientWebSocketExtSealed<Body> + Sized + Send + Sync + 'static
1407{
1408    /// Create a new [`WebSocketRequestBuilder`]] to be used to establish a WebSocket connection over http/1.1.
1409    fn websocket(&self, url: impl IntoUrl) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
1410
1411    /// Create a new [`WebSocketRequestBuilder`] to be used to establish a WebSocket connection over h2.
1412    fn websocket_h2(
1413        &self,
1414        url: impl IntoUrl,
1415    ) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
1416
1417    /// Create a new [`WebSocketRequestBuilder`] starting from the given request.
1418    ///
1419    /// This is useful in cases where you already have a request that you wish to use,
1420    /// for example in the case of a proxied reuqest.
1421    fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
1422        &self,
1423        req: Request<RequestBody>,
1424    ) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
1425}
1426
1427impl<S, Body> HttpClientWebSocketExt<Body> for S
1428where
1429    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
1430{
1431    fn websocket(&self, url: impl IntoUrl) -> WebSocketRequestBuilder<WithService<'_, Self, Body>> {
1432        WebSocketRequestBuilder::new_with_service(self, url)
1433    }
1434
1435    fn websocket_h2(
1436        &self,
1437        url: impl IntoUrl,
1438    ) -> WebSocketRequestBuilder<WithService<'_, Self, Body>> {
1439        WebSocketRequestBuilder::new_h2_with_service(self, url)
1440    }
1441
1442    fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
1443        &self,
1444        req: Request<RequestBody>,
1445    ) -> WebSocketRequestBuilder<WithService<'_, Self, Body>> {
1446        WebSocketRequestBuilder::new_with_service_and_request(self, req)
1447    }
1448}
1449
1450/// Extends a blocking HTTP client with WebSocket handshake builders.
1451///
1452/// The HTTP client remains reusable after the handshake. Each successful
1453/// handshake returns one independent, connected [`BlockingClientWebSocket`].
1454///
1455/// # Panics
1456///
1457/// Blocking handshakes and socket I/O must not run directly on an asynchronous
1458/// executor thread.
1459///
1460/// ```no_run
1461/// use rama_core::{Service, error::BoxError};
1462/// use rama_http::{
1463///     Body, Request, Response,
1464///     service::client::blocking::Client,
1465/// };
1466/// use rama_ws::handshake::client::BlockingHttpClientWebSocketExt as _;
1467///
1468/// fn exchange<S>(client: &Client<S>) -> Result<(), BoxError>
1469/// where
1470///     S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
1471/// {
1472///     let mut socket = client
1473///         .websocket("wss://example.com/chat")
1474///         .with_header("authorization", "Bearer secret")
1475///         .try_handshake()?;
1476///
1477///     socket.send_message("hello".into())?;
1478///     let _reply = socket.recv_message()?;
1479///     Ok(())
1480/// }
1481/// ```
1482pub trait BlockingHttpClientWebSocketExt<Body>:
1483    private::BlockingHttpClientWebSocketExtSealed<Body>
1484{
1485    /// The asynchronous service wrapped by this blocking HTTP client.
1486    type AsyncService: Service<Request, Output = Response<Body>, Error: Into<BoxError>>;
1487
1488    /// Create a WebSocket request builder for an HTTP/1.1 upgrade.
1489    fn websocket(
1490        &self,
1491        url: impl IntoUrl,
1492    ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
1493
1494    /// Create a WebSocket request builder for HTTP/2 Extended CONNECT.
1495    fn websocket_h2(
1496        &self,
1497        url: impl IntoUrl,
1498    ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
1499
1500    /// Create a WebSocket request builder from an existing request.
1501    fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
1502        &self,
1503        request: Request<RequestBody>,
1504    ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
1505}
1506
1507impl<S, Body> BlockingHttpClientWebSocketExt<Body> for BlockingHttpClient<S>
1508where
1509    S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
1510    Body: Send + 'static,
1511{
1512    type AsyncService = S;
1513
1514    fn websocket(
1515        &self,
1516        url: impl IntoUrl,
1517    ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body> {
1518        BlockingWebSocketRequestBuilder::new_blocking_with_service(self, url)
1519    }
1520
1521    fn websocket_h2(
1522        &self,
1523        url: impl IntoUrl,
1524    ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body> {
1525        BlockingWebSocketRequestBuilder::new_blocking_h2_with_service(self, url)
1526    }
1527
1528    fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
1529        &self,
1530        request: Request<RequestBody>,
1531    ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body> {
1532        BlockingWebSocketRequestBuilder::new_blocking_with_service_and_request(self, request)
1533    }
1534}
1535
1536mod private {
1537    use super::*;
1538
1539    pub trait HttpClientWebSocketExtSealed<Body> {}
1540
1541    impl<S, Body> HttpClientWebSocketExtSealed<Body> for S where
1542        S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>
1543    {
1544    }
1545
1546    pub trait BlockingHttpClientWebSocketExtSealed<Body> {}
1547
1548    impl<S, Body> BlockingHttpClientWebSocketExtSealed<Body> for BlockingHttpClient<S>
1549    where
1550        S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
1551        Body: Send + 'static,
1552    {
1553    }
1554}
1555
1556#[cfg(test)]
1557mod tests {
1558    use super::*;
1559    use rama_core::{ServiceInput, bytes::Bytes, service::service_fn};
1560    use rama_http::HeaderMap;
1561    use std::sync::{
1562        Arc,
1563        atomic::{AtomicUsize, Ordering},
1564    };
1565
1566    struct ResponseLease(Arc<AtomicUsize>);
1567
1568    impl Drop for ResponseLease {
1569        fn drop(&mut self) {
1570            self.0.fetch_add(1, Ordering::Release);
1571        }
1572    }
1573
1574    #[test]
1575    fn blocking_client_websocket_roundtrip_and_lifetimes() {
1576        fn assert_send<T: Send>() {}
1577        assert_send::<BlockingClientWebSocket>();
1578
1579        let leases_dropped = Arc::new(AtomicUsize::new(0));
1580        let service_leases_dropped = leases_dropped.clone();
1581        let service = service_fn(move |request: Request| {
1582            let leases_dropped = service_leases_dropped.clone();
1583            async move {
1584                let is_h2 = request.version() == Version::HTTP_2;
1585                let accept = if is_h2 {
1586                    assert_eq!(request.method(), Method::CONNECT);
1587                    assert!(request.headers().typed_get::<SecWebSocketKey>().is_none());
1588                    None
1589                } else {
1590                    assert_eq!(request.method(), Method::GET);
1591                    let key = request
1592                        .headers()
1593                        .typed_get::<SecWebSocketKey>()
1594                        .expect("HTTP/1.1 client handshake request to contain a key");
1595                    Some(
1596                        headers::SecWebSocketAccept::try_from(key)
1597                            .expect("client handshake key to produce an accept value"),
1598                    )
1599                };
1600
1601                if request.uri().path().is_some_and(|path| path == "/custom") {
1602                    assert_eq!(
1603                        request.headers().get("x-rama-test"),
1604                        Some(&rama_http::HeaderValue::from_static("custom")),
1605                    );
1606                }
1607
1608                let (client_io, server_io) = tokio::io::duplex(4 * 1024);
1609                let (pending, on_upgrade) = rama_http::io::upgrade::pending();
1610                pending.fulfill(rama_http::io::upgrade::Upgraded::new(
1611                    ServiceInput::new(client_io),
1612                    Bytes::new(),
1613                ));
1614
1615                tokio::spawn(async move {
1616                    let mut socket = AsyncWebSocket::from_raw_socket(
1617                        ServiceInput::new(server_io),
1618                        Role::Server,
1619                        None,
1620                    )
1621                    .await;
1622                    let message = socket.recv_message().await.unwrap();
1623                    socket.send_message(message).await.unwrap();
1624                });
1625
1626                let mut response = Response::new(ResponseLease(leases_dropped));
1627                if let Some(accept) = accept {
1628                    *response.status_mut() = StatusCode::SWITCHING_PROTOCOLS;
1629                    *response.version_mut() = Version::HTTP_11;
1630                    response
1631                        .headers_mut()
1632                        .typed_insert(headers::Upgrade::websocket());
1633                    response
1634                        .headers_mut()
1635                        .typed_insert(headers::Connection::upgrade());
1636                    response.headers_mut().typed_insert(accept);
1637                } else {
1638                    *response.status_mut() = StatusCode::OK;
1639                    *response.version_mut() = Version::HTTP_2;
1640                }
1641                response.extensions().insert(on_upgrade);
1642                Ok::<_, BoxError>(response)
1643            }
1644        });
1645
1646        let client = BlockingHttpClient::try_new(service).unwrap();
1647        let client_clone = client.clone();
1648        drop(client);
1649
1650        let config = WebSocketConfig::default().with_read_buffer_size(4 * 1024);
1651        let mut from_url = client_clone
1652            .websocket("ws://example.test/echo")
1653            .with_config(config)
1654            .try_handshake()
1655            .unwrap();
1656        assert_eq!(from_url.response().status, StatusCode::SWITCHING_PROTOCOLS);
1657        assert_eq!(from_url.get_config().read_buffer_size, 4 * 1024);
1658        assert_eq!(
1659            from_url
1660                .extensions()
1661                .get_ref::<StreamTransformed>()
1662                .unwrap()
1663                .by,
1664            "rama-http::Upgraded",
1665        );
1666
1667        let request = Request::builder()
1668            .version(Version::HTTP_11)
1669            .uri("ws://example.test/custom")
1670            .header("x-rama-test", "custom")
1671            .body(Body::empty())
1672            .unwrap();
1673        let mut from_request = client_clone
1674            .websocket_with_request(request)
1675            .try_handshake()
1676            .unwrap();
1677        let mut from_h2 = client_clone
1678            .websocket_h2("wss://example.test/h2")
1679            .try_handshake()
1680            .unwrap();
1681
1682        drop(client_clone);
1683        assert_eq!(leases_dropped.load(Ordering::Acquire), 0);
1684
1685        from_url.send_message("from url".into()).unwrap();
1686        assert_eq!(
1687            from_url
1688                .recv_message()
1689                .unwrap()
1690                .into_text()
1691                .unwrap()
1692                .as_str(),
1693            "from url",
1694        );
1695
1696        from_request.send_message("from request".into()).unwrap();
1697        assert_eq!(
1698            from_request
1699                .recv_message()
1700                .unwrap()
1701                .into_text()
1702                .unwrap()
1703                .as_str(),
1704            "from request",
1705        );
1706
1707        from_h2.send_message("from h2".into()).unwrap();
1708        assert_eq!(
1709            from_h2
1710                .recv_message()
1711                .unwrap()
1712                .into_text()
1713                .unwrap()
1714                .as_str(),
1715            "from h2",
1716        );
1717
1718        let BlockingClientWebSocket {
1719            socket: from_url,
1720            response,
1721            accepted_protocol: protocol,
1722        } = from_url;
1723        assert_eq!(response.status, StatusCode::SWITCHING_PROTOCOLS);
1724        assert!(protocol.is_none());
1725        assert_eq!(leases_dropped.load(Ordering::Acquire), 0);
1726        drop(from_url);
1727        assert_eq!(leases_dropped.load(Ordering::Acquire), 1);
1728        drop(from_request);
1729        assert_eq!(leases_dropped.load(Ordering::Acquire), 2);
1730        drop(from_h2);
1731        assert_eq!(leases_dropped.load(Ordering::Acquire), 3);
1732    }
1733
1734    #[cfg(feature = "dial9")]
1735    #[test]
1736    fn blocking_handshake_runs_inside_dial9_session() {
1737        let temp_dir = tempfile::tempdir().unwrap();
1738        let config = rama_core::telemetry::dial9::Dial9Config::builder()
1739            .enabled(true)
1740            .base_path(temp_dir.path().join("blocking-websocket.bin"))
1741            .max_file_size(1024 * 1024)
1742            .max_total_size(4 * 1024 * 1024)
1743            .build()
1744            .unwrap();
1745        let runtime = rama_core::rt::blocking::Runtime::builder()
1746            .with_dial9_config(config)
1747            .try_build()
1748            .unwrap();
1749        let service = service_fn(|request: Request| async move {
1750            assert!(
1751                rama_core::telemetry::dial9::telemetry::TelemetryHandle::current().is_enabled()
1752            );
1753            let key = request
1754                .headers()
1755                .typed_get::<SecWebSocketKey>()
1756                .expect("handshake request to contain a key");
1757            let (client_io, _server_io) = tokio::io::duplex(1024);
1758            let (pending, on_upgrade) = rama_http::io::upgrade::pending();
1759            pending.fulfill(rama_http::io::upgrade::Upgraded::new(
1760                ServiceInput::new(client_io),
1761                Bytes::new(),
1762            ));
1763
1764            let mut response = Response::new(());
1765            *response.status_mut() = StatusCode::SWITCHING_PROTOCOLS;
1766            *response.version_mut() = Version::HTTP_11;
1767            response
1768                .headers_mut()
1769                .typed_insert(headers::Upgrade::websocket());
1770            response
1771                .headers_mut()
1772                .typed_insert(headers::Connection::upgrade());
1773            response.headers_mut().typed_insert(
1774                headers::SecWebSocketAccept::try_from(key)
1775                    .expect("client handshake key to produce an accept value"),
1776            );
1777            response.extensions().insert(on_upgrade);
1778            Ok::<_, BoxError>(response)
1779        });
1780        let client = BlockingHttpClient::with_runtime(service, &runtime);
1781
1782        let socket = client
1783            .websocket("wss://example.test/socket")
1784            .try_handshake()
1785            .unwrap();
1786        assert_eq!(socket.response().status, StatusCode::SWITCHING_PROTOCOLS);
1787    }
1788
1789    fn offered_pmd(raw: &str) -> Option<SecWebSocketExtensions> {
1790        let mut headers = HeaderMap::new();
1791        headers.insert(
1792            header::SEC_WEBSOCKET_EXTENSIONS,
1793            raw.parse().expect("valid sec-websocket-extensions header"),
1794        );
1795        headers.typed_get::<SecWebSocketExtensions>()
1796    }
1797
1798    fn h2_response_with_pmd(raw: &str) -> Response<()> {
1799        let mut response = Response::new(());
1800        *response.version_mut() = Version::HTTP_2;
1801        *response.status_mut() = StatusCode::OK;
1802        response.headers_mut().insert(
1803            header::SEC_WEBSOCKET_EXTENSIONS,
1804            raw.parse().expect("valid sec-websocket-extensions header"),
1805        );
1806        response
1807    }
1808
1809    #[test]
1810    fn h2_handshake_accepts_any_successful_connect_status() {
1811        let mut response = Response::new(());
1812        *response.version_mut() = Version::HTTP_2;
1813        *response.status_mut() = StatusCode::CREATED;
1814
1815        validate_http_server_response(&response, None, None, None)
1816            .expect("successful CONNECT response");
1817
1818        *response.status_mut() = StatusCode::BAD_REQUEST;
1819        assert!(matches!(
1820            validate_http_server_response(&response, None, None, None),
1821            Err(ResponseValidateError::UnexpectedStatusCode(
1822                StatusCode::BAD_REQUEST
1823            ))
1824        ));
1825    }
1826
1827    /// Validate an (h2) server handshake response carrying `server_raw` against
1828    /// a client that offered `offered_raw`, returning the negotiated
1829    /// `client_max_window_bits`.
1830    fn validate_pmd(
1831        server_raw: &str,
1832        offered_raw: &str,
1833    ) -> Result<Option<u8>, ResponseValidateError> {
1834        let response = h2_response_with_pmd(server_raw);
1835        let accepted =
1836            validate_http_server_response(&response, None, None, offered_pmd(offered_raw))?;
1837        match accepted.extension {
1838            Some(Extension::PerMessageDeflate(cfg)) => Ok(cfg.client_max_window_bits),
1839            other => panic!("expected per-message-deflate extension, got {other:?}"),
1840        }
1841    }
1842
1843    // Regression: a valueless `client_max_window_bits` offer is parsed as the
1844    // sentinel `Some(0)` ("server may pick any value <= 15"). A server response
1845    // of `client_max_window_bits=15` must be accepted, not rejected as an
1846    // extension mismatch. Previously the `srv > offered` check evaluated
1847    // `15 > 0` and falsely failed the handshake (intermittent WS-over-h2 502s).
1848    #[test]
1849    fn valueless_client_max_window_bits_accepts_server_choice() {
1850        assert_eq!(
1851            Some(15),
1852            validate_pmd(
1853                "permessage-deflate; client_max_window_bits=15",
1854                "permessage-deflate; client_max_window_bits",
1855            )
1856            .expect("valueless offer should accept the server's window bits"),
1857        );
1858    }
1859
1860    #[test]
1861    fn explicit_client_max_window_bits_rejects_larger_server_choice() {
1862        assert!(matches!(
1863            validate_pmd(
1864                "permessage-deflate; client_max_window_bits=15",
1865                "permessage-deflate; client_max_window_bits=10",
1866            ),
1867            Err(ResponseValidateError::ExtensionMismatch(_)),
1868        ));
1869    }
1870
1871    #[test]
1872    fn explicit_client_max_window_bits_accepts_smaller_server_choice() {
1873        assert_eq!(
1874            Some(10),
1875            validate_pmd(
1876                "permessage-deflate; client_max_window_bits=10",
1877                "permessage-deflate; client_max_window_bits=12",
1878            )
1879            .expect("server choosing a smaller window should validate"),
1880        );
1881    }
1882}