Skip to main content

ntex/ws/
client.rs

1//! WebSocket client.
2use std::{fmt, marker, pin};
3
4#[cfg(feature = "openssl")]
5use crate::connect::openssl;
6#[cfg(feature = "openssl")]
7use tls_openssl::ssl::SslConnector;
8
9#[cfg(feature = "rustls")]
10use crate::connect::rustls::{TlsClientFilter, TlsConnector};
11#[cfg(feature = "rustls")]
12use tls_rustls::ClientConfig as RustlsClientConfig;
13
14use base64::{Engine, engine::general_purpose::STANDARD as base64};
15use nanorand::Rng;
16use urly::Url;
17
18use crate::client::{ClientCodec, ClientConfig, ClientRawRequest, ClientResponse, host_header};
19use crate::connect::{Connect, ConnectError, Connector};
20use crate::error::{Error, ErrorMapping};
21use crate::http::body::BodySize;
22use crate::http::header::{self, HeaderMap, HeaderValue};
23use crate::http::{ConnectionType, Message, Method, RequestHead, StatusCode};
24use crate::io::{Base, DispatchItem, Dispatcher, Filter, Io, Layer, Reason, Sealed};
25use crate::service::{IntoService, Pipeline, apply_fn, fn_service};
26use crate::util::{Either, select};
27use crate::{Cfg, Service, SharedCfg, channel::mpsc, rt, time::timeout, ws};
28
29use super::cfg::is_token;
30use super::error::{WsClientError, WsConfigError, WsError};
31use super::proto::{CloseCode, CloseReason};
32use super::{WsClientConfig, handshake::header_contains_token, transport::WsTransport};
33
34thread_local! {
35    static CFG: SharedCfg = SharedCfg::new("WS-CLIENT").into();
36}
37
38/// Builder for establishing a WebSocket client connection.
39///
40/// The builder contains the target URI and a typed [`WsClientConfig`]. Use
41/// [`connect`](Self::connect) to perform the opening handshake.
42pub struct WsClient<F> {
43    uri: Url,
44    err: Option<WsConfigError>,
45    cfg: Cfg<WsClientConfig>,
46    http_cfg: Cfg<ClientConfig>,
47    connector: Pipeline<Connect<Url>, Io<F>, Error<ConnectError>>,
48    filter: marker::PhantomData<F>,
49}
50
51impl WsClient<Base> {
52    /// Creates a client for `uri` using the supplied configuration.
53    ///
54    /// ```rust
55    /// use ntex::{SharedCfg, time::Seconds};
56    /// use ntex::ws::{WsClient, WsClientConfig};
57    ///
58    /// #[ntex::main]
59    /// async fn main() {
60    ///     let cfg = SharedCfg::new("WS-CLIENT").add(
61    ///         WsClientConfig::new()
62    ///             .set_max_frame_size(128 * 1024)
63    ///             .set_handshake_timeout(Seconds(10))
64    ///     );
65    ///
66    ///     let _client = WsClient::new("ws://localhost/socket", cfg);
67    /// }
68    /// ```
69    ///
70    /// URI conversion and validation errors are stored and returned by
71    /// [`connect`](Self::connect).
72    pub fn new<U>(uri: U, cfg: impl Into<Cfg<WsClientConfig>>) -> Self
73    where
74        Url: TryFrom<U>,
75        WsConfigError: From<<Url as TryFrom<U>>::Error>,
76    {
77        let (uri, err) = match Url::try_from(uri) {
78            Ok(uri) => {
79                let err = match uri.scheme_str() {
80                    _ if uri.host().is_none() => Some(WsConfigError::MissingHost),
81                    Some("http" | "ws" | "https" | "wss") => None,
82                    Some(_) => Some(WsConfigError::UnknownScheme),
83                    None => Some(WsConfigError::MissingScheme),
84                };
85                (uri, err)
86            }
87            Err(err) => (Url::new(), Some(WsConfigError::from(err))),
88        };
89
90        let cfg = cfg.into();
91        let shared = cfg.shared();
92
93        WsClient {
94            uri,
95            err,
96            cfg,
97            http_cfg: shared.get(),
98            connector: Pipeline::new(shared, Connector::<Url>::new()),
99            filter: marker::PhantomData,
100        }
101    }
102}
103
104impl<F> WsClient<F> {
105    /// Replaces the network connector used to establish the connection.
106    pub fn connector<U, S>(self, f: impl IntoService<S, SharedCfg, Connect<Url>>) -> WsClient<U>
107    where
108        U: Filter + 'static,
109        S: Service<SharedCfg, Connect<Url>, Res = Io<U>, Error = Error<ConnectError>> + 'static,
110    {
111        let shared = self.cfg.shared();
112        WsClient {
113            uri: self.uri,
114            err: self.err,
115            cfg: self.cfg,
116            http_cfg: self.http_cfg,
117            connector: Pipeline::new(shared, f.into_service()),
118            filter: marker::PhantomData,
119        }
120    }
121
122    #[cfg(feature = "openssl")]
123    /// Uses the supplied OpenSSL connector for secure connections.
124    pub fn openssl(self, config: SslConnector) -> WsClient<Layer<openssl::SslFilter>> {
125        self.connector(openssl::SslConnector::new(config))
126    }
127
128    #[cfg(feature = "rustls")]
129    /// Uses the supplied rustls connector for secure connections.
130    pub fn rustls(
131        self,
132        config: std::sync::Arc<RustlsClientConfig>,
133    ) -> WsClient<Layer<TlsClientFilter>> {
134        self.connector(TlsConnector::from(config))
135    }
136}
137
138impl<F> WsClient<F>
139where
140    F: Filter,
141{
142    /// Establishes the connection and performs the WebSocket opening handshake.
143    ///
144    /// # Errors
145    ///
146    /// Returns an error if connection establishment, HTTP encoding or decoding,
147    /// URI validation, timeout handling, or handshake validation fails.
148    pub async fn connect(&self) -> Result<WsConnection<F>, Error<WsClientError>> {
149        if let Some(err) = self.err.clone() {
150            return Err(Error::from(WsClientError::Config(err)).with_service(self.cfg.service()));
151        }
152
153        let mut head = self.request_head();
154
155        // Generate a random key for the `Sec-WebSocket-Key` header.
156        // a base64-encoded (see Section 4 of [RFC4648]) value that,
157        // when decoded, is 16 bytes in length (RFC 6455)
158        let mut sec_key: [u8; 16] = [0; 16];
159        nanorand::tls_rng().fill(&mut sec_key);
160        let key = base64.encode(sec_key);
161
162        head.headers.insert(
163            header::SEC_WEBSOCKET_KEY,
164            HeaderValue::try_from(key.as_str()).unwrap(),
165        );
166
167        let msg = Connect::new(self.uri.clone()).set_addr(self.cfg.addr);
168        log::trace!(
169            "{}: Open ws connection to {:?} addr: {:?}",
170            self.cfg.tag(),
171            self.uri,
172            self.cfg.addr
173        );
174
175        // the connector attributes its own errors
176        let io = self.connector.call(msg).await.into_error()?;
177        self.handshake(io, head, &key)
178            .await
179            .map_err(|e| e.with_service(self.cfg.service()))
180    }
181
182    /// Sends the handshake request and validates the response.
183    async fn handshake(
184        &self,
185        io: Io<F>,
186        head: Message<RequestHead>,
187        key: &str,
188    ) -> Result<WsConnection<F>, Error<WsClientError>> {
189        let tag = io.tag();
190
191        // create Framed and send request
192        let codec = ClientCodec::new(true, io.shared().get());
193
194        // send request and read response
195        let fut = async {
196            log::trace!("{tag}: Sending ws handshake http message");
197            io.send(
198                ClientRawRequest {
199                    head,
200                    headers: None,
201                    size: BodySize::None,
202                }
203                .into(),
204                &codec,
205            )
206            .await?;
207            log::trace!("{tag}: Waiting for ws handshake response");
208            io.recv(&codec)
209                .await?
210                .ok_or(WsClientError::Disconnected(None))
211        };
212
213        // set request timeout
214        let response = if self.cfg.timeout.non_zero() {
215            timeout(self.cfg.timeout, fut)
216                .await
217                .map_err(|()| WsClientError::Timeout)
218                .and_then(|res| res)?
219        } else {
220            fut.await?
221        };
222        log::trace!("{tag}: Ws handshake response is received {response:?}");
223
224        // verify response
225        if response.status != StatusCode::SWITCHING_PROTOCOLS {
226            return Err(Error::from(WsClientError::InvalidResponseStatus(
227                response.status,
228            )));
229        }
230
231        // Check for "UPGRADE" to websocket header
232        if !header_contains_token(&response.headers, &header::UPGRADE, "websocket") {
233            log::trace!("{tag}: Invalid upgrade header");
234            return Err(Error::from(WsClientError::InvalidUpgradeHeader));
235        }
236
237        // Check for "CONNECTION" header
238        if let Some(conn) = response.headers.get(&header::CONNECTION) {
239            if !header_contains_token(&response.headers, &header::CONNECTION, "upgrade") {
240                log::trace!("{tag}: Invalid connection header: {conn:?}");
241                return Err(Error::from(WsClientError::InvalidConnectionHeader(
242                    conn.clone(),
243                )));
244            }
245        } else {
246            log::trace!("{tag}: Missing connection header");
247            return Err(Error::from(WsClientError::MissingConnectionHeader));
248        }
249
250        if let Some(hdr_key) = response.headers.get(&header::SEC_WEBSOCKET_ACCEPT) {
251            let encoded = ws::hash_key(key.as_ref()).map_err(|_| {
252                Error::from(WsClientError::InvalidChallengeResponse(
253                    String::new(),
254                    hdr_key.clone(),
255                ))
256            })?;
257            if hdr_key.as_bytes() != encoded.as_bytes() {
258                log::trace!(
259                    "{tag}: Invalid challenge response: expected: {encoded} received: {hdr_key:?}"
260                );
261                return Err(Error::from(WsClientError::InvalidChallengeResponse(
262                    encoded,
263                    hdr_key.clone(),
264                )));
265            }
266        } else {
267            log::trace!("{tag}: Missing SEC-WEBSOCKET-ACCEPT header");
268            return Err(Error::from(WsClientError::MissingWebSocketAcceptHeader));
269        }
270
271        validate_negotiation(&response.headers, &self.cfg.headers).map_err(Error::from)?;
272        log::trace!("{tag}: Ws handshake response verification is completed");
273
274        // response and ws io
275        Ok(WsConnection::new(
276            io,
277            ClientResponse::with_empty_payload(response, self.http_cfg.clone()),
278            if self.cfg.server_mode {
279                ws::Codec::new().max_size(self.cfg.max_size)
280            } else {
281                ws::Codec::new()
282                    .max_size(self.cfg.max_size)
283                    .set_client_mode()
284            },
285        ))
286    }
287}
288
289impl<F> WsClient<F> {
290    /// Creates the handshake request head without the `Sec-WebSocket-Key`.
291    fn request_head(&self) -> Message<RequestHead> {
292        let mut head = Message::<RequestHead>::new();
293        // the message pool may return a recycled head whose method is not GET
294        // (e.g. previously used by the HTTP/1 server dispatcher for a POST request)
295        head.method = Method::GET;
296        head.uri = self.uri.clone();
297        head.set_connection_type(ConnectionType::Upgrade);
298
299        // copy headers, the head is empty
300        for (key, value) in &self.cfg.headers {
301            head.headers_mut().append(key.clone(), value.clone());
302        }
303
304        // host header, without userinfo and the scheme's default port
305        if !head.headers.contains_key(header::HOST)
306            && let Some(val) = host_header(&self.uri)
307        {
308            head.headers.insert(header::HOST, val);
309        }
310
311        #[cfg(feature = "cookie")]
312        {
313            // set cookies, appended to a configured `Cookie` header
314            if let Some(ref jar) = self.cfg.cookies {
315                let mut cookie = Vec::new();
316                for value in head.headers.get_all(header::COOKIE) {
317                    if !cookie.is_empty() {
318                        cookie.extend_from_slice(b"; ");
319                    }
320                    cookie.extend_from_slice(value.as_bytes());
321                }
322                for c in jar.iter() {
323                    crate::http::helpers::push_cookie(&mut cookie, c.name(), c.value());
324                }
325                if let Ok(val) = HeaderValue::from_bytes(&cookie) {
326                    head.headers.insert(header::COOKIE, val);
327                }
328            }
329        }
330
331        head
332    }
333}
334
335fn validate_negotiation(response: &HeaderMap, offered: &HeaderMap) -> Result<(), WsClientError> {
336    if let Some(extensions) = response.get(header::SEC_WEBSOCKET_EXTENSIONS) {
337        return Err(WsClientError::UnexpectedWebSocketExtensions(
338            extensions.clone(),
339        ));
340    }
341
342    let mut protocols = response.get_all(header::SEC_WEBSOCKET_PROTOCOL);
343    if let Some(protocol) = protocols.next() {
344        let selected = protocol.to_str().ok();
345        let valid = protocols.next().is_none()
346            && selected.is_some_and(|selected| {
347                is_token(selected)
348                    && offered
349                        .get(header::SEC_WEBSOCKET_PROTOCOL)
350                        .and_then(|offered| offered.to_str().ok())
351                        .is_some_and(|offered| {
352                            offered.split(',').any(|item| item.trim() == selected)
353                        })
354            });
355        if !valid {
356            return Err(WsClientError::InvalidWebSocketProtocol(protocol.clone()));
357        }
358    }
359    Ok(())
360}
361
362impl<F> fmt::Debug for WsClient<F> {
363    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
364        f.debug_struct("WsClient").field("cfg", &self.cfg).finish()
365    }
366}
367
368/// An established WebSocket client connection.
369///
370/// This value retains the opening-handshake response, the WebSocket codec, and
371/// the underlying I/O stream.
372pub struct WsConnection<F> {
373    io: Io<F>,
374    sink: ws::WsSink,
375    res: ClientResponse,
376}
377
378impl<F> WsConnection<F> {
379    fn new(io: Io<F>, res: ClientResponse, codec: ws::Codec) -> Self {
380        // the sink is also the dispatcher's codec, they share the codec state
381        let sink = ws::WsSink::new(io.get_ref(), codec, io.shared().get());
382        Self { io, sink, res }
383    }
384
385    /// Returns the connection's WebSocket codec.
386    pub fn codec(&self) -> &ws::Codec {
387        self.sink.codec()
388    }
389
390    /// Returns the opening-handshake response.
391    pub fn response(&self) -> &ClientResponse {
392        &self.res
393    }
394}
395
396impl<F> WsConnection<F> {
397    /// Returns a sink for sending messages over this connection.
398    ///
399    /// All sinks of a connection share the same state, so a message cannot be
400    /// sent through any of them once one has sent a close message.
401    pub fn sink(&self) -> ws::WsSink {
402        self.sink.clone()
403    }
404
405    /// Consumes the connection and returns its I/O stream, codec, and
406    /// opening-handshake response.
407    pub fn into_inner(self) -> (Io<F>, ws::Codec, ClientResponse) {
408        (self.io, self.sink.codec().clone(), self.res)
409    }
410}
411
412impl WsConnection<Sealed> {
413    /// Starts the WebSocket dispatcher and returns a channel of received frames.
414    ///
415    /// The dispatcher runs in a spawned task. Protocol and connection errors
416    /// are delivered through the returned channel. A close frame from the peer
417    /// is answered automatically, unless a close message has already been sent
418    /// through a sink of this connection. Dropping the receiver sends
419    /// a close frame and closes the connection once the peer responds or the
420    /// closing-handshake timeout expires.
421    pub fn receiver(self) -> mpsc::Receiver<Result<ws::Frame, WsError<()>>> {
422        let (tx, rx): (_, mpsc::Receiver<Result<ws::Frame, WsError<()>>>) = mpsc::channel();
423
424        rt::spawn(async move {
425            let tx2 = tx.clone();
426            let io = self.io.get_ref();
427            let sink = self.sink();
428            let sink2 = sink.clone();
429
430            let fut = self.start(fn_service(async move |item: ws::Frame| {
431                if let ws::Frame::Close(reason) = &item
432                    && !sink2.is_closed()
433                {
434                    // answer the peer's close frame, echoing its code
435                    let reply = reason.as_ref().map(|r| CloseReason::from(r.code));
436                    if sink2.send(ws::Message::Close(reply)).await.is_err() {
437                        let reply = CloseReason::from(CloseCode::Normal);
438                        let _ = sink2.send(ws::Message::Close(Some(reply))).await;
439                    }
440                }
441                match tx.send(Ok(item)) {
442                    Ok(()) => (),
443                    Err(_) => io.close(),
444                }
445                Ok::<Option<ws::Message>, ()>(None)
446            }));
447            let mut fut = pin::pin!(fut);
448
449            let result = match select(fut.as_mut(), tx2.closed()).await {
450                Either::Left(result) => result,
451                Either::Right(()) => {
452                    // the receiver is dropped, start the closing handshake
453                    let _ = sink
454                        .send(ws::Message::Close(Some(CloseCode::Normal.into())))
455                        .await;
456                    fut.await
457                }
458            };
459
460            if let Err(e) = result {
461                let _ = tx2.send(Err(e));
462            }
463        });
464
465        rx
466    }
467
468    /// Runs the WebSocket dispatcher with `svc` handling received frames.
469    ///
470    /// The service may return a message to send to the peer or [`None`] when no
471    /// response is required.
472    pub async fn start<T>(
473        self,
474        svc: impl IntoService<T, (), ws::Frame>,
475    ) -> Result<(), WsError<T::Error>>
476    where
477        T: Service<(), ws::Frame, Res = Option<ws::Message>> + 'static,
478    {
479        let io = self.io.get_ref();
480        let sink = self.sink();
481        let service = apply_fn(
482            svc.into_service().map_err(WsError::Service),
483            async move |req, svc| match req {
484                DispatchItem::<ws::WsSink>::Item(item) => {
485                    let close = matches!(item, ws::Frame::Close(_));
486                    let result = svc.call(item).await;
487                    if matches!(&result, Ok(Some(ws::Message::Close(_)))) {
488                        sink.start_close_timeout();
489                    }
490                    if close {
491                        let io = io.clone();
492                        rt::spawn(async move { io.close() });
493                    }
494                    result
495                }
496                // a clean disconnect is not an error
497                DispatchItem::Control(_) | DispatchItem::Stop(Reason::Io(None)) => Ok(None),
498                DispatchItem::Stop(Reason::Service) => {
499                    Ok(Some(ws::Message::Close(Some(CloseReason {
500                        code: CloseCode::Away,
501                        description: None,
502                    }))))
503                }
504                DispatchItem::Stop(Reason::KeepAlive) => Err(WsError::KeepAlive),
505                DispatchItem::Stop(Reason::ReadTimeout) => Err(WsError::ReadTimeout),
506                DispatchItem::Stop(Reason::WriteTimeout) => Err(WsError::WriteTimeout),
507                DispatchItem::Stop(Reason::Decoder(e)) => {
508                    if !sink.is_closed() {
509                        let reason = CloseReason::from(CloseCode::Protocol);
510                        let _ = sink.send(ws::Message::Close(Some(reason))).await;
511                    }
512                    Err(WsError::Protocol(e))
513                }
514                DispatchItem::Stop(Reason::Encoder(e)) => Err(WsError::Protocol(e)),
515                DispatchItem::Stop(Reason::Io(e)) => Err(WsError::Disconnected(e)),
516            },
517        );
518
519        Dispatcher::new(self.io, self.sink, Pipeline::new((), service)).await
520    }
521}
522
523impl<F: Filter> WsConnection<F> {
524    /// Erases the concrete I/O filter type.
525    pub fn seal(self) -> WsConnection<Sealed> {
526        WsConnection {
527            io: self.io.seal(),
528            sink: self.sink,
529            res: self.res,
530        }
531    }
532
533    /// Converts the connection into a binary WebSocket transport.
534    pub fn into_transport(self) -> Io<Layer<WsTransport, F>> {
535        WsTransport::create(self.io, self.sink.codec().clone())
536    }
537}
538
539impl<F> fmt::Debug for WsConnection<F> {
540    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
541        f.debug_struct("WsConnection")
542            .field("response", &self.res)
543            .finish()
544    }
545}
546
547#[cfg(test)]
548mod tests {
549    use super::*;
550
551    #[crate::rt_test]
552    async fn test_debug() {
553        let client = WsClient::new("http://localhost", SharedCfg::default());
554        assert!(format!("{client:?}").contains("WsClient"));
555    }
556
557    #[crate::rt_test]
558    async fn request_head_keeps_all_header_values() {
559        let mut cfg = WsClientConfig::new();
560        cfg.headers
561            .append(header::ACCEPT, HeaderValue::from_static("a"));
562        cfg.headers
563            .append(header::ACCEPT, HeaderValue::from_static("b"));
564        #[cfg(feature = "cookie")]
565        {
566            cfg.headers
567                .append(header::COOKIE, HeaderValue::from_static("x=1"));
568            cfg.headers
569                .append(header::COOKIE, HeaderValue::from_static("y=2"));
570            cfg = cfg.set_cookie(coo_kie::Cookie::new("z", "3"));
571        }
572        let client = WsClient::new("http://localhost", SharedCfg::new("WS").add(cfg));
573
574        let head = client.request_head();
575        let values: Vec<_> = head.headers.get_all(header::ACCEPT).collect();
576        assert_eq!(values, ["a", "b"]);
577        #[cfg(feature = "cookie")]
578        assert_eq!(head.headers.get(header::COOKIE).unwrap(), "x=1; y=2; z=3");
579    }
580
581    #[crate::rt_test]
582    async fn header_override() {
583        let cfg = WsClientConfig::new()
584            .set_header(header::CONTENT_TYPE, "111")
585            .unwrap()
586            .set_header(header::CONTENT_TYPE, "222")
587            .unwrap();
588
589        assert_eq!(
590            cfg.headers
591                .get(header::CONTENT_TYPE)
592                .unwrap()
593                .to_str()
594                .unwrap(),
595            "222"
596        );
597    }
598
599    #[test]
600    fn protocols() {
601        let cfg = WsClientConfig::new()
602            .set_protocols(["chat", "superchat"])
603            .unwrap();
604        assert_eq!(
605            cfg.headers
606                .get(header::SEC_WEBSOCKET_PROTOCOL)
607                .unwrap()
608                .to_str()
609                .unwrap(),
610            "chat,superchat"
611        );
612
613        let cfg = cfg.set_protocols([] as [&str; 0]).unwrap();
614        assert!(!cfg.headers.contains_key(header::SEC_WEBSOCKET_PROTOCOL));
615        assert!(WsClientConfig::new().set_protocols(["bad\n"]).is_err());
616        assert!(
617            WsClientConfig::new()
618                .set_protocols(["bad protocol"])
619                .is_err()
620        );
621        assert!(
622            WsClientConfig::new()
623                .set_protocols(["first,second"])
624                .is_err()
625        );
626    }
627
628    #[test]
629    fn negotiation() {
630        let configured = WsClientConfig::new()
631            .set_protocols(["chat", "superchat"])
632            .unwrap();
633        let mut response = HeaderMap::new();
634
635        response.insert(
636            header::SEC_WEBSOCKET_PROTOCOL,
637            HeaderValue::from_static("chat"),
638        );
639        validate_negotiation(&response, &configured.headers).unwrap();
640
641        let mut offered_headers = HeaderMap::new();
642        offered_headers.insert(
643            header::SEC_WEBSOCKET_PROTOCOL,
644            HeaderValue::from_static("chat, superchat"),
645        );
646        response.insert(
647            header::SEC_WEBSOCKET_PROTOCOL,
648            HeaderValue::from_static("superchat"),
649        );
650        validate_negotiation(&response, &offered_headers).unwrap();
651
652        response.insert(
653            header::SEC_WEBSOCKET_PROTOCOL,
654            HeaderValue::from_static("other"),
655        );
656        assert!(matches!(
657            validate_negotiation(&response, &configured.headers),
658            Err(WsClientError::InvalidWebSocketProtocol(_))
659        ));
660
661        response.insert(
662            header::SEC_WEBSOCKET_PROTOCOL,
663            HeaderValue::from_static("chat,superchat"),
664        );
665        assert!(matches!(
666            validate_negotiation(&response, &configured.headers),
667            Err(WsClientError::InvalidWebSocketProtocol(_))
668        ));
669
670        response.remove(header::SEC_WEBSOCKET_PROTOCOL);
671        response.append(
672            header::SEC_WEBSOCKET_PROTOCOL,
673            HeaderValue::from_static("chat"),
674        );
675        response.append(
676            header::SEC_WEBSOCKET_PROTOCOL,
677            HeaderValue::from_static("superchat"),
678        );
679        assert!(matches!(
680            validate_negotiation(&response, &configured.headers),
681            Err(WsClientError::InvalidWebSocketProtocol(_))
682        ));
683
684        response.remove(header::SEC_WEBSOCKET_PROTOCOL);
685        response.insert(
686            header::SEC_WEBSOCKET_EXTENSIONS,
687            HeaderValue::from_static("permessage-deflate"),
688        );
689        assert!(matches!(
690            validate_negotiation(&response, &configured.headers),
691            Err(WsClientError::UnexpectedWebSocketExtensions(_))
692        ));
693    }
694
695    #[crate::rt_test]
696    async fn basic_errs() {
697        let err = WsClient::new("//localhost", SharedCfg::default())
698            .connect()
699            .await
700            .err()
701            .unwrap();
702        assert!(matches!(
703            err.into_error(),
704            WsClientError::Config(WsConfigError::MissingScheme)
705        ));
706
707        let err = WsClient::new("unknown://localhost", SharedCfg::default())
708            .connect()
709            .await
710            .err()
711            .unwrap();
712        assert!(matches!(
713            err.into_error(),
714            WsClientError::Config(WsConfigError::UnknownScheme)
715        ));
716
717        let err = WsClient::new("/", SharedCfg::default())
718            .connect()
719            .await
720            .err()
721            .unwrap();
722        assert!(matches!(
723            err.into_error(),
724            WsClientError::Config(WsConfigError::MissingHost)
725        ));
726    }
727
728    #[crate::rt_test]
729    async fn basic_auth() {
730        let cfg = WsClientConfig::new()
731            .set_basic_auth("username", Some("password"))
732            .unwrap();
733        assert_eq!(
734            cfg.headers
735                .get(header::AUTHORIZATION)
736                .unwrap()
737                .to_str()
738                .unwrap(),
739            "Basic dXNlcm5hbWU6cGFzc3dvcmQ="
740        );
741
742        let cfg = WsClientConfig::new()
743            .set_basic_auth("username", None)
744            .unwrap();
745        assert_eq!(
746            cfg.headers
747                .get(header::AUTHORIZATION)
748                .unwrap()
749                .to_str()
750                .unwrap(),
751            "Basic dXNlcm5hbWU6"
752        );
753
754        let cfg = cfg.set_basic_auth("username", Some("password")).unwrap();
755        assert_eq!(
756            cfg.headers
757                .get(header::AUTHORIZATION)
758                .unwrap()
759                .to_str()
760                .unwrap(),
761            "Basic dXNlcm5hbWU6cGFzc3dvcmQ="
762        );
763    }
764
765    #[crate::rt_test]
766    async fn bearer_auth() {
767        let cfg = WsClientConfig::new()
768            .set_bearer_auth("someS3cr3tAutht0k3n")
769            .unwrap();
770        assert_eq!(
771            cfg.headers
772                .get(header::AUTHORIZATION)
773                .unwrap()
774                .to_str()
775                .unwrap(),
776            "Bearer someS3cr3tAutht0k3n"
777        );
778    }
779
780    #[cfg(feature = "cookie")]
781    #[crate::rt_test]
782    async fn basics() {
783        use coo_kie::Cookie;
784
785        let cfg = WsClientConfig::new()
786            .set_origin("test-origin")
787            .unwrap()
788            .set_max_frame_size(100)
789            .set_server_mode()
790            .set_protocols(["v1", "v2"])
791            .unwrap()
792            .set_header_if_none(header::CONTENT_TYPE, "json")
793            .unwrap()
794            .set_header_if_none(header::CONTENT_TYPE, "text")
795            .unwrap()
796            .set_cookie(Cookie::build(("cookie1", "value1")));
797
798        assert!(cfg.server_mode);
799        assert_eq!(cfg.max_size, 100);
800
801        assert!(WsClient::new("/", SharedCfg::default()).err.is_some());
802        assert!(
803            WsClient::new("http:///test", SharedCfg::default())
804                .err
805                .is_some()
806        );
807        assert!(
808            WsClient::new("hmm://test.com/", SharedCfg::default())
809                .err
810                .is_some()
811        );
812    }
813
814    /// Runs `connect()` over an in-memory stream, returns the handshake request.
815    async fn handshake_request(uri: &str, cfg: WsClientConfig) -> String {
816        use crate::{testing::IoTest, util::Bytes};
817        use std::cell::RefCell;
818
819        let (client, server) = IoTest::create();
820        client.remote_buffer_cap(4096);
821        let io = RefCell::new(Some(Io::new(server, SharedCfg::default())));
822        let ws = WsClient::new(uri, cfg).connector(fn_service(async move |_: Connect<Url>| {
823            Ok::<_, Error<ConnectError>>(io.borrow_mut().take().unwrap())
824        }));
825        let fut = rt::spawn(async move { ws.connect().await.map(drop) });
826
827        let mut req = Vec::new();
828        while !req.ends_with(b"\r\n\r\n") {
829            let buf: Bytes = client.read().await.unwrap();
830            req.extend_from_slice(&buf);
831        }
832        client.close().await;
833        let _ = fut.await;
834        String::from_utf8(req).unwrap()
835    }
836
837    #[crate::rt_test]
838    async fn pooled_request_head_method_is_get() {
839        // a request head released back to the thread-local message pool keeps its
840        // method (e.g. POST from the HTTP/1 server dispatcher); a ws client built
841        // from such a recycled head must still send a GET handshake
842        let mut head = Message::<RequestHead>::new();
843        head.method = Method::POST;
844        drop(head);
845
846        let req = handshake_request("ws://localhost/", WsClientConfig::new()).await;
847        assert!(req.starts_with("GET / HTTP/1.1\r\n"), "{req}");
848    }
849
850    #[cfg(feature = "cookie")]
851    #[crate::rt_test]
852    async fn cookies_extend_configured_header() {
853        use coo_kie::Cookie;
854
855        let cfg = || {
856            WsClientConfig::new()
857                .set_cookie(Cookie::build(("c1", "v1")))
858                .set_cookie(Cookie::build(("c2", "v2")))
859        };
860        let req = handshake_request("ws://localhost/", cfg()).await;
861        let cookie = req
862            .lines()
863            .find_map(|l| l.strip_prefix("cookie: "))
864            .unwrap();
865        let mut cookies: Vec<_> = cookie.split("; ").collect();
866        cookies.sort_unstable();
867        assert_eq!(cookies, ["c1=v1", "c2=v2"]);
868
869        let cfg = cfg().set_header(header::COOKIE, "c0=v0").unwrap();
870        let req = handshake_request("ws://localhost/", cfg).await;
871        let cookie = req
872            .lines()
873            .find_map(|l| l.strip_prefix("cookie: "))
874            .unwrap();
875        assert!(cookie.starts_with("c0=v0; "), "{cookie}");
876        let mut cookies: Vec<_> = cookie.split("; ").collect();
877        cookies.sort_unstable();
878        assert_eq!(cookies, ["c0=v0", "c1=v1", "c2=v2"]);
879    }
880
881    type Connected = (
882        Result<WsConnection<Base>, Error<WsClientError>>,
883        crate::testing::IoTest,
884    );
885
886    /// Runs `connect()` against an in-memory peer that answers the handshake
887    /// with the response produced by `response`.
888    async fn connect_with(
889        cfg: WsClientConfig,
890        io_cfg: SharedCfg,
891        response: impl FnOnce(String) -> String,
892    ) -> Connected {
893        use crate::{testing::IoTest, util::Bytes};
894        use std::cell::RefCell;
895
896        let (client, server) = IoTest::create();
897        client.remote_buffer_cap(4096);
898        let io = RefCell::new(Some(Io::new(server, io_cfg)));
899        let ws = WsClient::new("ws://localhost/", SharedCfg::new("WS").add(cfg)).connector(
900            fn_service(async move |_: Connect<Url>| {
901                Ok::<_, Error<ConnectError>>(io.borrow_mut().take().unwrap())
902            }),
903        );
904        let fut = rt::spawn(async move { ws.connect().await });
905
906        let mut req = Vec::new();
907        while !req.ends_with(b"\r\n\r\n") {
908            let buf: Bytes = client.read().await.unwrap();
909            req.extend_from_slice(&buf);
910        }
911        let req = String::from_utf8(req).unwrap();
912        let key = req
913            .lines()
914            .find_map(|l| l.strip_prefix("sec-websocket-key: "))
915            .unwrap();
916        let accept = ws::hash_key(key.as_bytes()).unwrap();
917        client.write(response(accept));
918        (fut.await.unwrap(), client)
919    }
920
921    fn switching(headers: &str) -> String {
922        format!("HTTP/1.1 101 Switching Protocols\r\n{headers}\r\n")
923    }
924
925    fn valid(accept: &str) -> String {
926        switching(&format!(
927            "upgrade: websocket\r\nconnection: upgrade\r\nsec-websocket-accept: {accept}\r\n"
928        ))
929    }
930
931    async fn connected(cfg: WsClientConfig, io_cfg: SharedCfg) -> Connected {
932        connect_with(cfg, io_cfg, |accept| valid(&accept)).await
933    }
934
935    #[crate::rt_test]
936    async fn handshake_response_errors() {
937        async fn err(response: impl FnOnce(String) -> String) -> WsClientError {
938            let cfg = WsClientConfig::new().set_handshake_timeout(0);
939            let (res, _client) = connect_with(cfg, SharedCfg::default(), response).await;
940            res.unwrap_err().into_error()
941        }
942
943        assert!(matches!(
944            err(|_| switching("upgrade: h2c\r\nconnection: upgrade\r\n")).await,
945            WsClientError::InvalidUpgradeHeader
946        ));
947        assert!(matches!(
948            err(|_| switching("upgrade: websocket\r\nconnection: close\r\n")).await,
949            WsClientError::InvalidConnectionHeader(val) if val == "close"
950        ));
951        assert!(matches!(
952            err(|_| switching("upgrade: websocket\r\n")).await,
953            WsClientError::MissingConnectionHeader
954        ));
955        assert!(matches!(
956            err(|_| switching("upgrade: websocket\r\nconnection: upgrade\r\n")).await,
957            WsClientError::MissingWebSocketAcceptHeader
958        ));
959        assert!(matches!(
960            err(|_| valid("aW52YWxpZA==")).await,
961            WsClientError::InvalidChallengeResponse(_, val) if val == "aW52YWxpZA=="
962        ));
963    }
964
965    fn peer_frame(codec: &ws::Codec, msg: ws::Message) -> crate::util::Bytes {
966        let mut dst = crate::util::BytePages::default();
967        crate::codec::Encoder::encode(codec, msg, &mut dst).unwrap();
968        dst.into()
969    }
970
971    fn read_frame(client: &crate::testing::IoTest, codec: &ws::Codec) -> ws::Frame {
972        let mut data = crate::util::BytesMut::from(&client.read_any()[..]);
973        crate::codec::Decoder::decode(codec, &mut data)
974            .unwrap()
975            .unwrap()
976    }
977
978    #[crate::rt_test]
979    async fn server_mode_connection() {
980        let cfg = WsClientConfig::new().set_server_mode();
981        let (res, client) = connected(cfg, SharedCfg::default()).await;
982        let conn = res.unwrap();
983        assert!(format!("{conn:?}").contains("WsConnection"));
984        assert_eq!(conn.response().status(), StatusCode::SWITCHING_PROTOCOLS);
985        assert!(!conn.codec().is_closed());
986
987        // servers cannot send 1010, the peer's close is answered with 1000
988        let rx = conn.seal().receiver();
989        client.write(peer_frame(
990            &ws::Codec::new().set_client_mode(),
991            ws::Message::Close(Some(CloseCode::Extension.into())),
992        ));
993        let item = rx.recv().await.unwrap().unwrap();
994        assert_eq!(item, ws::Frame::Close(Some(CloseCode::Extension.into())));
995        crate::time::sleep(crate::time::Millis(50)).await;
996        assert_eq!(
997            read_frame(&client, &ws::Codec::new().set_client_mode()),
998            ws::Frame::Close(Some(CloseCode::Normal.into()))
999        );
1000    }
1001
1002    #[crate::rt_test]
1003    async fn start_service_error_sends_away_close() {
1004        let (res, client) = connected(WsClientConfig::new(), SharedCfg::default()).await;
1005        let conn = res.unwrap().seal();
1006
1007        client.write(peer_frame(
1008            &ws::Codec::new(),
1009            ws::Message::Text("text".into()),
1010        ));
1011        let err = conn
1012            .start(fn_service(async |_: ws::Frame| {
1013                Err::<Option<ws::Message>, _>("err")
1014            }))
1015            .await
1016            .unwrap_err();
1017        assert!(matches!(err, WsError::Service("err")));
1018        assert_eq!(
1019            read_frame(&client, &ws::Codec::new()),
1020            ws::Frame::Close(Some(CloseCode::Away.into()))
1021        );
1022    }
1023
1024    #[crate::rt_test]
1025    async fn start_encoder_error() {
1026        let (res, client) = connected(WsClientConfig::new(), SharedCfg::default()).await;
1027        let conn = res.unwrap().seal();
1028
1029        client.write(peer_frame(
1030            &ws::Codec::new(),
1031            ws::Message::Text("text".into()),
1032        ));
1033        let err = conn
1034            .start(fn_service(async |_: ws::Frame| {
1035                Ok::<_, ()>(Some(ws::Message::Ping(vec![0; 126].into())))
1036            }))
1037            .await
1038            .unwrap_err();
1039        assert!(matches!(
1040            err,
1041            WsError::Protocol(ws::error::ProtocolError::InvalidLength(126))
1042        ));
1043    }
1044
1045    #[crate::rt_test]
1046    async fn start_io_error() {
1047        let (res, client) = connected(WsClientConfig::new(), SharedCfg::default()).await;
1048        let conn = res.unwrap().seal();
1049
1050        client.read_error(std::io::Error::other("failed"));
1051        let err = conn
1052            .start(fn_service(async |_: ws::Frame| Ok::<_, ()>(None)))
1053            .await
1054            .unwrap_err();
1055        assert!(matches!(err, WsError::Disconnected(Some(_))));
1056    }
1057
1058    #[crate::rt_test]
1059    async fn start_keepalive() {
1060        let io_cfg = SharedCfg::new("KA")
1061            .add(crate::io::IoConfig::new().set_keepalive_timeout(crate::time::Seconds(1)));
1062        let (res, _client) = connected(WsClientConfig::new(), io_cfg.into()).await;
1063        let err = res
1064            .unwrap()
1065            .seal()
1066            .start(fn_service(async |_: ws::Frame| Ok::<_, ()>(None)))
1067            .await
1068            .unwrap_err();
1069        assert!(matches!(err, WsError::KeepAlive));
1070    }
1071}