1#![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#[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)]
52pub 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
95pub 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
117pub mod websocket_builder_mode {
119 use std::sync::Arc;
120
121 use rama_core::rt::blocking::Runtime;
122
123 #[derive(Debug)]
125 #[non_exhaustive]
126 pub struct Async;
127
128 #[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
137pub 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 _ => (),
217 }
218 service.build_from_request(request)
219}
220
221#[derive(Debug)]
222pub 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)]
234pub 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
302pub 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 let response_status = response.status();
323 if response_status != StatusCode::SWITCHING_PROTOCOLS {
324 return Err(ResponseValidateError::UnexpectedStatusCode(response_status));
325 }
326
327 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 !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 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 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 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 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 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 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 #[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 #[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 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 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 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 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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 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 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 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 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 pub fn try_handshake(self) -> Result<BlockingClientWebSocket, HandshakeError> {
1011 self.try_handshake_with_extensions(Extensions::new())
1012 }
1013
1014 #[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 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 pub fn key(mut self, key: Option<headers::SecWebSocketKey>) -> Self {
1059 self.key = key;
1060 self
1061 }
1062 }
1063}
1064
1065pub 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
1117pub 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 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)]
1220pub struct ClientWebSocket<S = AsyncWebSocket> {
1224 pub socket: S,
1226 pub response: response::Parts,
1228 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 #[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 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 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 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 pub fn response(&self) -> &response::Parts {
1334 &self.response
1335 }
1336
1337 pub fn accepted_protocol(&self) -> Option<&str> {
1339 self.accepted_protocol.as_ref().map(|p| p.0.as_ref())
1340 }
1341
1342 pub fn into_inner(self) -> S {
1344 self.socket
1345 }
1346}
1347
1348pub type BlockingWebSocket = WebSocket<BlockingIo<rama_http::io::upgrade::Upgraded>>;
1351
1352#[derive(Debug)]
1354pub struct BlockingClientWebSocket {
1355 pub socket: BlockingWebSocket,
1357 pub response: response::Parts,
1359 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 pub fn response(&self) -> &response::Parts {
1380 &self.response
1381 }
1382
1383 pub fn accepted_protocol(&self) -> Option<&str> {
1385 self.accepted_protocol.as_ref().map(|p| p.0.as_ref())
1386 }
1387
1388 pub fn send_message(&mut self, message: Message) -> Result<(), ProtocolError> {
1390 self.socket.send(message)
1391 }
1392
1393 pub fn recv_message(&mut self) -> Result<Message, ProtocolError> {
1395 self.socket.read()
1396 }
1397
1398 pub fn into_inner(self) -> BlockingWebSocket {
1400 self.socket
1401 }
1402}
1403
1404pub trait HttpClientWebSocketExt<Body>:
1406 private::HttpClientWebSocketExtSealed<Body> + Sized + Send + Sync + 'static
1407{
1408 fn websocket(&self, url: impl IntoUrl) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
1410
1411 fn websocket_h2(
1413 &self,
1414 url: impl IntoUrl,
1415 ) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
1416
1417 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
1450pub trait BlockingHttpClientWebSocketExt<Body>:
1483 private::BlockingHttpClientWebSocketExtSealed<Body>
1484{
1485 type AsyncService: Service<Request, Output = Response<Body>, Error: Into<BoxError>>;
1487
1488 fn websocket(
1490 &self,
1491 url: impl IntoUrl,
1492 ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
1493
1494 fn websocket_h2(
1496 &self,
1497 url: impl IntoUrl,
1498 ) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
1499
1500 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 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 #[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}