Skip to main content

slim_config/websocket/
client.rs

1// Copyright AGNTCY Contributors (https://github.com/agntcy)
2// SPDX-License-Identifier: Apache-2.0
3
4use std::net::SocketAddr;
5use std::str::FromStr;
6use std::sync::Arc;
7use std::time::Duration;
8
9use base64::Engine;
10use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
11use bytes::Bytes;
12use fastwebsockets::{Role, WebSocket, handshake};
13use http::header::{CONNECTION, HOST, ORIGIN, UPGRADE, USER_AGENT};
14use http::header::{HeaderName, HeaderValue};
15use http::uri::Authority;
16use http::{Request, Response, StatusCode};
17use http_body_util::Empty;
18use hyper::body::Incoming;
19use hyper_util::client::legacy::Client;
20use hyper_util::client::legacy::connect::HttpInfo;
21use hyper_util::rt::{TokioExecutor, TokioIo};
22use sha1::{Digest, Sha1};
23use tower::{ServiceBuilder, ServiceExt};
24use tracing::warn;
25
26use crate::auth::ClientAuthenticator;
27use crate::client::{AuthenticationConfig as ClientAuthConfig, ClientConfig};
28use crate::errors::ConfigError;
29use crate::transport::TransportProtocol;
30use crate::transport_common::{Alpn, ProxyTunnel, build_https_connector};
31
32use super::common::{UpgradedWebSocket, WebSocketEndpoint};
33
34/// The error type emitted by hyper-util's legacy Client. Auth and bearer→query
35/// layers preserve `S::Error`, so the whole stack produces this single type.
36type ClientError = hyper_util::client::legacy::Error;
37
38#[derive(Clone)]
39pub struct WebSocketClientChannel {
40    inner: Arc<WebSocketClientChannelInner>,
41}
42
43struct WebSocketClientChannelInner {
44    /// The upgraded websocket. Wrapped in `Mutex<Option<_>>` because the
45    /// underlying `fastwebsockets::WebSocket` can only be `.split()` once;
46    /// callers consume it via [`WebSocketClientChannel::take_websocket`].
47    websocket: parking_lot::Mutex<Option<UpgradedWebSocket>>,
48    local_addr: Option<SocketAddr>,
49    remote_addr: Option<SocketAddr>,
50}
51
52impl WebSocketClientChannel {
53    pub(crate) fn new(
54        websocket: UpgradedWebSocket,
55        local_addr: Option<SocketAddr>,
56        remote_addr: Option<SocketAddr>,
57    ) -> Self {
58        Self {
59            inner: Arc::new(WebSocketClientChannelInner {
60                websocket: parking_lot::Mutex::new(Some(websocket)),
61                local_addr,
62                remote_addr,
63            }),
64        }
65    }
66
67    /// Take ownership of the underlying upgraded websocket. Returns `None`
68    /// if it has already been taken — the websocket can only be split into
69    /// read/write halves once.
70    pub fn take_websocket(&self) -> Option<UpgradedWebSocket> {
71        self.inner.websocket.lock().take()
72    }
73
74    pub fn local_addr(&self) -> Option<SocketAddr> {
75        self.inner.local_addr
76    }
77
78    pub fn remote_addr(&self) -> Option<SocketAddr> {
79        self.inner.remote_addr
80    }
81}
82
83impl ClientConfig {
84    /// Build a WebSocket channel. Crate-private; external callers should use
85    /// [`ClientConfig::to_channel`].
86    pub(crate) async fn to_websocket_channel(&self) -> Result<WebSocketClientChannel, ConfigError> {
87        if self.resolved_transport() != TransportProtocol::Websocket {
88            return Err(ConfigError::WebSocketClientUnsupportedTransport);
89        }
90
91        let endpoint = WebSocketEndpoint::parse(self.endpoint.as_str())?;
92
93        self.retry_connect(|| self.connect_websocket_once(&endpoint))
94            .await
95    }
96
97    async fn connect_websocket_once(
98        &self,
99        endpoint: &WebSocketEndpoint,
100    ) -> Result<WebSocketClientChannel, ConfigError> {
101        // Build the upgrade request. We use http:// / https:// scheme since
102        // hyper rejects ws:// / wss://.
103        let request_uri = endpoint.http_request_uri()?;
104        let request = build_handshake_request(self, endpoint, request_uri)?;
105        let sent_key = request
106            .headers()
107            .get("Sec-WebSocket-Key")
108            .cloned()
109            .expect("Sec-WebSocket-Key set by build_handshake_request");
110
111        // Single budget covering connector setup + TLS + proxy CONNECT + HTTP
112        // request. The retry layer above counts the *attempt*, not each leg.
113        let total_timeout: Duration = self.connect_timeout.into();
114        let mut response: Response<Incoming> = self
115            .build_and_send(endpoint, request, total_timeout)
116            .await?;
117
118        if response.status() != StatusCode::SWITCHING_PROTOCOLS {
119            return Err(ConfigError::WebSocketHandshakeStatus(response.status()));
120        }
121        verify_sec_websocket_accept(&response, sent_key.as_bytes())?;
122
123        // HttpInfo is attached to the response by hyper-util's legacy Client.
124        // Pull remote_addr off before we consume the response via upgrade::on.
125        let remote_addr = response
126            .extensions()
127            .get::<HttpInfo>()
128            .map(|info: &HttpInfo| info.remote_addr());
129
130        let upgraded = hyper::upgrade::on(&mut response)
131            .await
132            .map_err(|e| ConfigError::WebSocketConnection(std::io::Error::other(e)))?;
133
134        let websocket = WebSocket::after_handshake(TokioIo::new(upgraded), Role::Client);
135
136        Ok(WebSocketClientChannel::new(
137            websocket,
138            // hyper-util's HttpInfo only exposes remote_addr().
139            None,
140            remote_addr,
141        ))
142    }
143
144    /// Build the connector chain, layer auth + bearer→query on top, and send
145    /// the upgrade request. Each `match` arm calls into the generic
146    /// [`Self::send_via`] / [`Self::send_authed`] so the concrete connector
147    /// type is monomorphized end-to-end — no boxing required.
148    async fn build_and_send(
149        &self,
150        endpoint: &WebSocketEndpoint,
151        request: Request<Empty<Bytes>>,
152        timeout: Duration,
153    ) -> Result<Response<Incoming>, ConfigError> {
154        let base = self.create_http_connector()?;
155        let tls = if endpoint.secure {
156            self.load_tls_config().await?
157        } else {
158            None
159        };
160        if endpoint.secure && tls.is_none() {
161            return Err(ConfigError::WebSocketServerTlsMissing);
162        }
163
164        // hyper-util's proxy `Matcher` needs an http/https URI; the WS one
165        // built with ws:// / wss:// schemes is rejected.
166        let target_for_proxy = format!(
167            "{}://{}{}",
168            if endpoint.secure { "https" } else { "http" },
169            endpoint.authority,
170            endpoint.path,
171        );
172        let proxy_hit = self.proxy.should_use_proxy(&target_for_proxy);
173
174        let server_name = self
175            .server_name
176            .clone()
177            .or_else(|| self.origin.as_deref().and_then(host_from_authority))
178            .unwrap_or_else(|| endpoint.host.clone());
179
180        match (proxy_hit, tls) {
181            (None, None) => self.send_via(build_client(base), request, timeout).await,
182            (None, Some(tls)) => {
183                let https = build_https_connector(base, tls, Some(server_name), Alpn::Http1);
184                self.send_via(build_client(https), request, timeout).await
185            }
186            (Some(intercept), None) => match self
187                .build_proxy_tunnel(intercept, base, Alpn::Http1)
188                .await?
189            {
190                ProxyTunnel::Http(t) => self.send_via(build_client(t), request, timeout).await,
191                ProxyTunnel::Https(t) => self.send_via(build_client(t), request, timeout).await,
192            },
193            (Some(intercept), Some(tls)) => match self
194                .build_proxy_tunnel(intercept, base, Alpn::Http1)
195                .await?
196            {
197                ProxyTunnel::Http(t) => {
198                    let https = build_https_connector(t, tls, Some(server_name), Alpn::Http1);
199                    self.send_via(build_client(https), request, timeout).await
200                }
201                ProxyTunnel::Https(t) => {
202                    let https = build_https_connector(t, tls, Some(server_name), Alpn::Http1);
203                    self.send_via(build_client(https), request, timeout).await
204                }
205            },
206        }
207    }
208
209    /// Initialize the auth layer (if any), wrap `client` with it, and dispatch
210    /// the request. Generic over the connector type so no type erasure is
211    /// needed.
212    async fn send_via<C>(
213        &self,
214        client: Client<C, Empty<Bytes>>,
215        request: Request<Empty<Bytes>>,
216        timeout: Duration,
217    ) -> Result<Response<Incoming>, ConfigError>
218    where
219        C: hyper_util::client::legacy::connect::Connect + Clone + Send + Sync + 'static,
220    {
221        match &self.auth {
222            ClientAuthConfig::None => run_send(client.oneshot(request), timeout).await,
223            ClientAuthConfig::Basic(basic) => {
224                let layer = basic.get_client_layer()?;
225                self.warn_insecure_auth();
226                let svc = ServiceBuilder::new().layer(layer).service(client);
227                run_send(svc.oneshot(request), timeout).await
228            }
229            ClientAuthConfig::StaticJwt(jwt) => {
230                let mut layer = jwt.get_client_layer()?;
231                layer.initialize().await?;
232                self.warn_insecure_auth();
233                let svc = ServiceBuilder::new().layer(layer).service(client);
234                run_send(svc.oneshot(request), timeout).await
235            }
236            ClientAuthConfig::Jwt(jwt) => {
237                let mut layer = jwt.get_client_layer()?;
238                layer.initialize().await?;
239                self.warn_insecure_auth();
240                let svc = ServiceBuilder::new().layer(layer).service(client);
241                run_send(svc.oneshot(request), timeout).await
242            }
243            ClientAuthConfig::Oidc(cfg) => {
244                let mut layer = cfg.get_client_layer()?;
245                layer.initialize().await?;
246                self.warn_insecure_auth();
247                let svc = ServiceBuilder::new().layer(layer).service(client);
248                run_send(svc.oneshot(request), timeout).await
249            }
250            #[cfg(not(target_family = "windows"))]
251            ClientAuthConfig::Spire(spire) => {
252                let mut layer = spire.get_client_layer()?;
253                layer.initialize().await?;
254                self.warn_insecure_auth();
255                let svc = ServiceBuilder::new().layer(layer).service(client);
256                run_send(svc.oneshot(request), timeout).await
257            }
258        }
259    }
260}
261
262async fn run_send<F>(send: F, timeout: Duration) -> Result<Response<Incoming>, ConfigError>
263where
264    F: std::future::Future<Output = Result<Response<Incoming>, ClientError>>,
265{
266    if timeout.is_zero() {
267        send.await.map_err(box_send_err)
268    } else {
269        match tokio::time::timeout(timeout, send).await {
270            Ok(Ok(r)) => Ok(r),
271            Ok(Err(e)) => Err(box_send_err(e)),
272            Err(_) => Err(ConfigError::WebSocketHandshakeTimeout),
273        }
274    }
275}
276
277fn build_client<C>(connector: C) -> Client<C, Empty<Bytes>>
278where
279    C: hyper_util::client::legacy::connect::Connect + Clone + Send + Sync + 'static,
280{
281    Client::builder(TokioExecutor::new())
282        .pool_idle_timeout(Duration::from_secs(0))
283        .pool_max_idle_per_host(0)
284        .build(connector)
285}
286
287fn box_send_err<E>(err: E) -> ConfigError
288where
289    E: std::error::Error + Send + Sync + 'static,
290{
291    ConfigError::WebSocketClientSend(Box::new(err))
292}
293
294fn build_handshake_request(
295    config: &ClientConfig,
296    endpoint: &WebSocketEndpoint,
297    uri: http::Uri,
298) -> Result<Request<Empty<Bytes>>, ConfigError> {
299    const DEFAULT_USER_AGENT: &str = concat!("slim-websocket/", env!("CARGO_PKG_VERSION"));
300
301    // RFC 7230 §5.4: Host MUST be the URI's authority with userinfo (and the
302    // '@' delimiter) removed, and otherwise byte-identical.
303    let host_header = strip_userinfo(&endpoint.authority);
304
305    let mut request = Request::builder()
306        .method("GET")
307        .uri(uri)
308        .header(HOST, host_header)
309        .header(UPGRADE, "websocket")
310        .header(CONNECTION, "Upgrade")
311        .header("Sec-WebSocket-Key", handshake::generate_key())
312        .header("Sec-WebSocket-Version", "13")
313        .header(USER_AGENT, DEFAULT_USER_AGENT)
314        .body(Empty::<Bytes>::new())
315        .map_err(ConfigError::WebSocketRequest)?;
316
317    let headers = request.headers_mut();
318
319    if let Some(origin) = config.origin.as_deref() {
320        headers.insert(ORIGIN, HeaderValue::from_str(origin)?);
321    }
322
323    // Iterate config.headers in sorted order: HashMap iteration is
324    // non-deterministic, which would otherwise make request bytes vary
325    // between attempts (and break any snapshot-style tests).
326    let mut extra: Vec<(&String, &String)> = config.headers.iter().collect();
327    extra.sort_by(|a, b| a.0.cmp(b.0));
328    for (name, value) in extra {
329        let header_name = HeaderName::from_str(name)?;
330        // Filter reserved headers
331        if is_reserved_handshake_header(&header_name) {
332            warn!(
333                header = %header_name,
334                "ignoring reserved websocket handshake header supplied via config.headers",
335            );
336            continue;
337        }
338        headers.insert(header_name, HeaderValue::from_str(value)?);
339    }
340
341    Ok(request)
342}
343
344fn is_reserved_handshake_header(name: &HeaderName) -> bool {
345    matches!(
346        name.as_str(),
347        "host"
348            | "upgrade"
349            | "connection"
350            | "sec-websocket-key"
351            | "sec-websocket-version"
352            | "sec-websocket-accept"
353            | "sec-websocket-extensions"
354            | "sec-websocket-protocol"
355    )
356}
357
358/// Return `authority` with any leading `userinfo@` removed. Preserves the
359/// rest of the string verbatim (including absence-of-port), which is what
360/// RFC 7230 §5.4 requires for the Host header and what host-based ingress
361/// routing rules expect.
362fn strip_userinfo(authority: &str) -> &str {
363    match authority.rfind('@') {
364        Some(pos) => &authority[pos + 1..],
365        None => authority,
366    }
367}
368
369/// Strip any `:port` suffix from an HTTP authority, returning just the host.
370///
371/// IPv6 hosts in URI authorities are bracketed (`[::1]:8080`); delegating to
372/// [`http::uri::Authority`] handles that. Used to derive an SNI name from
373/// the `Origin` header when configured.
374fn host_from_authority(authority: &str) -> Option<String> {
375    if authority.is_empty() {
376        return None;
377    }
378    let parsed = Authority::from_str(authority).ok()?;
379    let host = parsed.host();
380    if host.is_empty() {
381        return None;
382    }
383    let unbracketed = host
384        .strip_prefix('[')
385        .and_then(|s| s.strip_suffix(']'))
386        .unwrap_or(host);
387    Some(unbracketed.to_string())
388}
389
390/// RFC 6455 §4.1: the server must echo `base64(sha1(client_key + GUID))` in
391/// `Sec-WebSocket-Accept`. `fastwebsockets::WebSocket::after_handshake`
392/// trusts the caller to have validated this, so we do it here.
393fn verify_sec_websocket_accept<B>(
394    response: &Response<B>,
395    sent_key: &[u8],
396) -> Result<(), ConfigError> {
397    const GUID: &[u8] = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
398
399    let header = response
400        .headers()
401        .get("Sec-WebSocket-Accept")
402        .ok_or(ConfigError::WebSocketMissingAcceptHeader)?;
403    let received = header
404        .to_str()
405        .map_err(|_| ConfigError::WebSocketAcceptMismatch)?;
406
407    let mut hasher = Sha1::new();
408    hasher.update(sent_key);
409    hasher.update(GUID);
410    let expected = BASE64_STANDARD.encode(hasher.finalize());
411
412    if expected == received {
413        Ok(())
414    } else {
415        Err(ConfigError::WebSocketAcceptMismatch)
416    }
417}
418
419// =====================================================================
420// Tests
421// =====================================================================
422#[cfg(test)]
423mod tests {
424    use super::*;
425
426    use crate::tls::client::TlsClientConfig;
427    use std::net::TcpListener;
428    use std::time::Duration;
429
430    fn available_port() -> u16 {
431        TcpListener::bind("127.0.0.1:0")
432            .expect("bind")
433            .local_addr()
434            .expect("local_addr")
435            .port()
436    }
437
438    #[tokio::test]
439    async fn test_websocket_client_invalid_endpoint_scheme() {
440        let cfg = ClientConfig::with_endpoint("http://127.0.0.1:80")
441            .with_tls_setting(TlsClientConfig::insecure());
442        let result = cfg.to_websocket_channel().await;
443        assert!(result.is_err(), "non-ws scheme must be rejected");
444    }
445
446    #[tokio::test]
447    async fn test_websocket_client_connect_refused() {
448        // Bind to grab a port, then drop the listener to guarantee it's closed.
449        let port = available_port();
450        let cfg = ClientConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
451            .with_tls_setting(TlsClientConfig::insecure())
452            .with_backoff(crate::client::BackoffConfig::new_fixed_interval(
453                Duration::from_millis(0),
454                1,
455            ));
456
457        let result = cfg.to_websocket_channel().await;
458        assert!(result.is_err(), "connection to closed port should fail");
459    }
460
461    #[tokio::test]
462    #[allow(clippy::disallowed_methods)]
463    async fn test_websocket_client_connect_timeout() {
464        let _env_guard = crate::test_env::PROXY_ENV_LOCK.lock().await;
465
466        // Isolate from shell/parallel tests that set HTTP_PROXY (see grpc::proxy tests).
467        let saved_proxy_env = [
468            ("http_proxy", std::env::var("http_proxy").ok()),
469            ("HTTP_PROXY", std::env::var("HTTP_PROXY").ok()),
470            ("https_proxy", std::env::var("https_proxy").ok()),
471            ("HTTPS_PROXY", std::env::var("HTTPS_PROXY").ok()),
472            ("all_proxy", std::env::var("all_proxy").ok()),
473            ("ALL_PROXY", std::env::var("ALL_PROXY").ok()),
474        ];
475        for (key, _) in &saved_proxy_env {
476            unsafe {
477                std::env::remove_var(key);
478            }
479        }
480
481        // RFC 5737 TEST-NET-1: guaranteed unroutable.
482        let cfg = ClientConfig::with_endpoint("ws://192.0.2.1:9")
483            .with_tls_setting(TlsClientConfig::insecure())
484            .with_connect_timeout(Duration::from_millis(200))
485            .with_backoff(crate::client::BackoffConfig::new_fixed_interval(
486                Duration::from_millis(0),
487                0,
488            ));
489
490        let start = std::time::Instant::now();
491        let outer = tokio::time::timeout(Duration::from_secs(2), cfg.to_websocket_channel()).await;
492        let elapsed = start.elapsed();
493
494        assert!(outer.is_ok(), "configured connect_timeout was not honored");
495        assert!(outer.unwrap().is_err(), "unroutable connect must fail");
496        assert!(
497            elapsed < Duration::from_secs(1),
498            "connect_timeout was not honored (took {elapsed:?})"
499        );
500
501        for (key, original) in saved_proxy_env {
502            unsafe {
503                match original {
504                    Some(value) => std::env::set_var(key, value),
505                    None => std::env::remove_var(key),
506                }
507            }
508        }
509    }
510
511    #[test]
512    fn test_host_from_authority_strips_port() {
513        assert_eq!(
514            host_from_authority("example.com:8443"),
515            Some("example.com".to_string())
516        );
517    }
518
519    #[test]
520    fn test_host_from_authority_no_port() {
521        assert_eq!(
522            host_from_authority("example.com"),
523            Some("example.com".to_string())
524        );
525    }
526
527    #[test]
528    fn test_host_from_authority_empty() {
529        assert_eq!(host_from_authority(""), None);
530    }
531
532    #[test]
533    fn test_host_from_authority_ipv6_with_port() {
534        // IPv6 literal: brackets must not be split as a port separator.
535        assert_eq!(
536            host_from_authority("[::1]:8080"),
537            Some("::1".to_string()),
538            "IPv6 host must be extracted intact"
539        );
540    }
541
542    #[test]
543    fn test_host_from_authority_ipv6_without_port() {
544        assert_eq!(
545            host_from_authority("[2001:db8::1]"),
546            Some("2001:db8::1".to_string()),
547        );
548    }
549
550    #[test]
551    fn test_is_reserved_handshake_header_matches_case_insensitively() {
552        for name in [
553            "host",
554            "Host",
555            "HOST",
556            "upgrade",
557            "connection",
558            "sec-websocket-key",
559            "Sec-WebSocket-Key",
560            "sec-websocket-version",
561            "sec-websocket-accept",
562            "sec-websocket-extensions",
563            "sec-websocket-protocol",
564        ] {
565            let h = HeaderName::from_str(name).expect("valid header name");
566            assert!(
567                is_reserved_handshake_header(&h),
568                "{name} should be reserved"
569            );
570        }
571    }
572
573    #[test]
574    fn test_is_reserved_handshake_header_allows_unrelated() {
575        for name in ["authorization", "origin", "x-trace-id", "user-agent"] {
576            let h = HeaderName::from_str(name).expect("valid header name");
577            assert!(
578                !is_reserved_handshake_header(&h),
579                "{name} must not be reserved"
580            );
581        }
582    }
583
584    fn handshake_request_for(cfg: &ClientConfig) -> Request<Empty<Bytes>> {
585        let endpoint = WebSocketEndpoint::parse(cfg.endpoint.as_str()).expect("endpoint");
586        let request_uri = endpoint.http_request_uri().expect("uri");
587        build_handshake_request(cfg, &endpoint, request_uri).expect("build")
588    }
589
590    #[test]
591    fn test_handshake_request_sets_required_headers() {
592        let cfg = ClientConfig::with_endpoint("ws://example.com:8080/p");
593        let req = handshake_request_for(&cfg);
594        let h = req.headers();
595        assert_eq!(h.get(HOST).unwrap(), "example.com:8080");
596        assert_eq!(h.get(UPGRADE).unwrap(), "websocket");
597        assert_eq!(h.get(CONNECTION).unwrap(), "Upgrade");
598        assert_eq!(h.get("Sec-WebSocket-Version").unwrap(), "13");
599        assert!(h.get("Sec-WebSocket-Key").is_some());
600        assert!(
601            h.get(USER_AGENT)
602                .and_then(|v| v.to_str().ok())
603                .map(|v| v.starts_with("slim-websocket/"))
604                .unwrap_or(false),
605            "User-Agent should default to slim-websocket/<version>"
606        );
607    }
608
609    #[test]
610    fn test_handshake_request_host_omits_default_ws_port() {
611        let cfg = ClientConfig::with_endpoint("ws://api.example.com/p");
612        let req = handshake_request_for(&cfg);
613        assert_eq!(req.headers().get(HOST).unwrap(), "api.example.com");
614    }
615
616    #[test]
617    fn test_handshake_request_host_omits_default_wss_port() {
618        let cfg = ClientConfig::with_endpoint("wss://api.example.com/p");
619        let req = handshake_request_for(&cfg);
620        assert_eq!(req.headers().get(HOST).unwrap(), "api.example.com");
621    }
622
623    #[test]
624    fn test_handshake_request_host_preserves_explicit_port() {
625        let cfg = ClientConfig::with_endpoint("ws://api.example.com:80/p");
626        let req = handshake_request_for(&cfg);
627        assert_eq!(req.headers().get(HOST).unwrap(), "api.example.com:80");
628    }
629
630    #[test]
631    fn test_handshake_request_host_excludes_userinfo() {
632        let cfg = ClientConfig::with_endpoint("ws://user:pw@example.com:8080/p");
633        let req = handshake_request_for(&cfg);
634        let host = req.headers().get(HOST).unwrap().to_str().unwrap();
635        assert!(
636            !host.contains('@'),
637            "Host header must not carry userinfo: {host}"
638        );
639        assert_eq!(host, "example.com:8080");
640    }
641
642    #[test]
643    fn test_handshake_request_host_for_ipv6() {
644        let cfg = ClientConfig::with_endpoint("ws://[::1]:9000/p");
645        let req = handshake_request_for(&cfg);
646        assert_eq!(req.headers().get(HOST).unwrap(), "[::1]:9000");
647    }
648
649    #[test]
650    fn test_strip_userinfo() {
651        assert_eq!(strip_userinfo("example.com:8080"), "example.com:8080");
652        assert_eq!(strip_userinfo("u@example.com"), "example.com");
653        assert_eq!(strip_userinfo("u:p@example.com:443"), "example.com:443");
654        assert_eq!(strip_userinfo("[::1]:9000"), "[::1]:9000");
655        assert_eq!(strip_userinfo("a@b@example.com"), "example.com");
656    }
657
658    #[test]
659    fn test_handshake_request_user_headers_cannot_clobber_reserved() {
660        let mut headers = std::collections::HashMap::new();
661        headers.insert("Upgrade".to_string(), "evil".to_string());
662        headers.insert("Connection".to_string(), "close".to_string());
663        headers.insert(
664            "Sec-WebSocket-Key".to_string(),
665            "AAAAAAAAAAAAAAAAAAAAAA==".to_string(),
666        );
667        headers.insert("Sec-WebSocket-Version".to_string(), "8".to_string());
668        headers.insert("Host".to_string(), "attacker.example".to_string());
669        let cfg = ClientConfig::with_endpoint("ws://example.com:8080/p").with_headers(headers);
670        let req = handshake_request_for(&cfg);
671        let h = req.headers();
672        assert_eq!(h.get(UPGRADE).unwrap(), "websocket");
673        assert_eq!(h.get(CONNECTION).unwrap(), "Upgrade");
674        assert_eq!(h.get("Sec-WebSocket-Version").unwrap(), "13");
675        assert_eq!(h.get(HOST).unwrap(), "example.com:8080");
676        let key = h.get("Sec-WebSocket-Key").unwrap().to_str().unwrap();
677        assert_ne!(key, "AAAAAAAAAAAAAAAAAAAAAA==");
678    }
679
680    #[test]
681    fn test_handshake_request_allows_user_supplied_non_reserved_headers() {
682        let mut headers = std::collections::HashMap::new();
683        headers.insert("X-Trace-Id".to_string(), "abc-123".to_string());
684        let cfg = ClientConfig::with_endpoint("ws://example.com:8080/p").with_headers(headers);
685        let req = handshake_request_for(&cfg);
686        assert_eq!(req.headers().get("X-Trace-Id").unwrap(), "abc-123");
687    }
688
689    #[test]
690    fn test_verify_sec_websocket_accept_matches_rfc_example() {
691        // RFC 6455 §1.3 worked example.
692        let key = "dGhlIHNhbXBsZSBub25jZQ==";
693        let expected = "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=";
694        let response: Response<()> = Response::builder()
695            .status(StatusCode::SWITCHING_PROTOCOLS)
696            .header("Sec-WebSocket-Accept", expected)
697            .body(())
698            .unwrap();
699        verify_sec_websocket_accept(&response, key.as_bytes()).expect("matches");
700    }
701
702    #[test]
703    fn test_verify_sec_websocket_accept_missing_header() {
704        let response: Response<()> = Response::builder()
705            .status(StatusCode::SWITCHING_PROTOCOLS)
706            .body(())
707            .unwrap();
708        assert!(matches!(
709            verify_sec_websocket_accept(&response, b"key"),
710            Err(ConfigError::WebSocketMissingAcceptHeader)
711        ));
712    }
713
714    #[test]
715    fn test_verify_sec_websocket_accept_wrong_digest() {
716        let response: Response<()> = Response::builder()
717            .status(StatusCode::SWITCHING_PROTOCOLS)
718            .header("Sec-WebSocket-Accept", "AAAAAAAAAAAAAAAAAAAAAAAAAAA=")
719            .body(())
720            .unwrap();
721        assert!(matches!(
722            verify_sec_websocket_accept(&response, b"dGhlIHNhbXBsZSBub25jZQ=="),
723            Err(ConfigError::WebSocketAcceptMismatch)
724        ));
725    }
726}