1use 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
38pub 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 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 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 pub fn openssl(self, config: SslConnector) -> WsClient<Layer<openssl::SslFilter>> {
125 self.connector(openssl::SslConnector::new(config))
126 }
127
128 #[cfg(feature = "rustls")]
129 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 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 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 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 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 let codec = ClientCodec::new(true, io.shared().get());
193
194 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 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 if response.status != StatusCode::SWITCHING_PROTOCOLS {
226 return Err(Error::from(WsClientError::InvalidResponseStatus(
227 response.status,
228 )));
229 }
230
231 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 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 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 fn request_head(&self) -> Message<RequestHead> {
292 let mut head = Message::<RequestHead>::new();
293 head.method = Method::GET;
296 head.uri = self.uri.clone();
297 head.set_connection_type(ConnectionType::Upgrade);
298
299 for (key, value) in &self.cfg.headers {
301 head.headers_mut().append(key.clone(), value.clone());
302 }
303
304 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 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
368pub 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 let sink = ws::WsSink::new(io.get_ref(), codec, io.shared().get());
382 Self { io, sink, res }
383 }
384
385 pub fn codec(&self) -> &ws::Codec {
387 self.sink.codec()
388 }
389
390 pub fn response(&self) -> &ClientResponse {
392 &self.res
393 }
394}
395
396impl<F> WsConnection<F> {
397 pub fn sink(&self) -> ws::WsSink {
402 self.sink.clone()
403 }
404
405 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 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 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 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 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 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 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 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 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 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 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 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}