Skip to main content

slim_config/websocket/
server.rs

1// Copyright AGNTCY Contributors (https://github.com/agntcy)
2// SPDX-License-Identifier: Apache-2.0
3
4use std::convert::Infallible;
5use std::future::Future;
6use std::net::SocketAddr;
7use std::pin::Pin;
8use std::sync::Arc;
9use std::sync::atomic::{AtomicBool, Ordering};
10use std::task::{Context, Poll};
11use std::time::Duration;
12
13use bytes::Bytes;
14use fastwebsockets::upgrade;
15use http_body_util::Empty;
16use hyper::Request;
17use hyper::Response;
18use hyper::StatusCode;
19use hyper::body::Incoming;
20use hyper::server::conn::http1;
21use hyper_util::rt::TokioIo;
22use hyper_util::service::TowerToHyperService;
23use slim_auth::jwt::VerifierJwt;
24use slim_auth::jwt_middleware::{PolicyCheckLayer, ValidateJwtLayer};
25use slim_auth::metadata::MetadataMap;
26use slim_auth::oidc::OidcVerifier;
27#[cfg(not(target_family = "windows"))]
28use slim_auth::spire::SpireIdentityManager;
29use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
30use tokio::net::{TcpListener, TcpStream};
31use tokio::sync::{OwnedSemaphorePermit, Semaphore};
32use tokio_rustls::TlsAcceptor;
33use tokio_util::sync::CancellationToken;
34use tower::util::BoxCloneService;
35use tower::{ServiceBuilder, service_fn};
36#[allow(deprecated)]
37use tower_http::auth::require_authorization::Basic;
38use tower_http::validate_request::ValidateRequestHeaderLayer;
39use tower_layer::Stack;
40use tracing::{debug, warn};
41
42use crate::auth::ServerAuthenticator;
43use crate::auth::jwt::Config as JwtAuthenticationConfig;
44use crate::auth::oidc::Config as OidcConfig;
45#[cfg(not(target_family = "windows"))]
46use crate::auth::spire::SpireConfig as SpireAuthConfig;
47use crate::errors::ConfigError;
48use crate::server::{AuthenticationConfig as ServerAuthConfig, ServerConfig};
49use crate::tls::common::RustlsConfigLoader;
50use crate::transport::TransportProtocol;
51use crate::websocket::query_token_layer::QueryTokenToAuthHeaderLayer;
52
53use super::common::{UpgradedWebSocket, WebSocketEndpoint};
54
55/// Maximum time allowed for the TLS handshake to complete after accepting a
56/// TCP connection. Prevents a silent or malicious client from pinning an
57/// accept task indefinitely.
58const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
59
60/// Maximum time allowed for a client to complete the HTTP request and
61/// WebSocket upgrade. Once the upgrade succeeds the connection is no longer
62/// bound by this timeout.
63const HTTP_UPGRADE_TIMEOUT: Duration = Duration::from_secs(10);
64
65/// Hard ceiling on the number of in-flight accept tasks when the
66/// configured `max_concurrent_streams` is `None`. Prevents an unbounded
67/// connection storm from exhausting memory / file descriptors on
68/// misconfigured servers.
69const DEFAULT_MAX_WEBSOCKET_CONNECTIONS: usize = 1024;
70
71/// Unified server-side stream: either a plain TCP stream or a TLS-wrapped
72/// TCP stream. Lets [`serve_connection`] be invoked with a single concrete
73/// type regardless of whether TLS is enabled.
74enum MaybeTlsStream {
75    Plain(TcpStream),
76    Tls(Box<tokio_rustls::server::TlsStream<TcpStream>>),
77}
78
79impl AsyncRead for MaybeTlsStream {
80    fn poll_read(
81        self: Pin<&mut Self>,
82        cx: &mut Context<'_>,
83        buf: &mut ReadBuf<'_>,
84    ) -> Poll<std::io::Result<()>> {
85        // Both inner types are `Unpin`, so projecting through `&mut *self`
86        // is safe without `pin-project`.
87        match self.get_mut() {
88            MaybeTlsStream::Plain(s) => Pin::new(s).poll_read(cx, buf),
89            MaybeTlsStream::Tls(s) => Pin::new(s.as_mut()).poll_read(cx, buf),
90        }
91    }
92}
93
94impl AsyncWrite for MaybeTlsStream {
95    fn poll_write(
96        self: Pin<&mut Self>,
97        cx: &mut Context<'_>,
98        buf: &[u8],
99    ) -> Poll<std::io::Result<usize>> {
100        match self.get_mut() {
101            MaybeTlsStream::Plain(s) => Pin::new(s).poll_write(cx, buf),
102            MaybeTlsStream::Tls(s) => Pin::new(s.as_mut()).poll_write(cx, buf),
103        }
104    }
105
106    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
107        match self.get_mut() {
108            MaybeTlsStream::Plain(s) => Pin::new(s).poll_flush(cx),
109            MaybeTlsStream::Tls(s) => Pin::new(s.as_mut()).poll_flush(cx),
110        }
111    }
112
113    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
114        match self.get_mut() {
115            MaybeTlsStream::Plain(s) => Pin::new(s).poll_shutdown(cx),
116            MaybeTlsStream::Tls(s) => Pin::new(s.as_mut()).poll_shutdown(cx),
117        }
118    }
119}
120
121pub struct AcceptedWebSocketConnection {
122    pub websocket: UpgradedWebSocket,
123    pub remote_addr: Option<SocketAddr>,
124    pub local_addr: Option<SocketAddr>,
125}
126
127pub type OnAcceptedWebSocket = Arc<
128    dyn Fn(AcceptedWebSocketConnection) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync,
129>;
130
131/// Resolved auth layers selected by [`ServerAuthConfig`]. The variant is
132/// built once at server startup and cloned into each accepted connection so
133/// the layer stack is composed at the point [`http1::Builder::serve_connection_with_upgrades`]
134/// is invoked. Each variant resolves to a different concrete tower
135/// `Service` type, which is why dispatch happens via `match` inside
136/// [`serve_connection`] rather than through a single boxed service.
137#[derive(Clone)]
138#[allow(clippy::large_enum_variant)]
139enum AuthKind {
140    None,
141    Basic(#[allow(deprecated)] ValidateRequestHeaderLayer<Basic<Empty<Bytes>>>),
142    Jwt(Stack<PolicyCheckLayer, ValidateJwtLayer<MetadataMap, VerifierJwt>>),
143    Oidc(Stack<PolicyCheckLayer, ValidateJwtLayer<MetadataMap, OidcVerifier>>),
144    #[cfg(not(target_family = "windows"))]
145    Spire(ValidateJwtLayer<MetadataMap, SpireIdentityManager>),
146}
147
148async fn build_auth_kind(config: &ServerConfig) -> Result<AuthKind, ConfigError> {
149    match &config.auth {
150        ServerAuthConfig::None => Ok(AuthKind::None),
151        ServerAuthConfig::Basic(basic) => {
152            let layer = <_ as ServerAuthenticator<Empty<Bytes>>>::get_server_layer(basic)?;
153            Ok(AuthKind::Basic(layer))
154        }
155        ServerAuthConfig::Jwt(jwt) => {
156            let layer = <JwtAuthenticationConfig as ServerAuthenticator<
157                Response<Empty<Bytes>>,
158            >>::get_server_layer(jwt)?;
159            Ok(AuthKind::Jwt(layer))
160        }
161        ServerAuthConfig::Oidc(oidc) => {
162            let layer =
163                <OidcConfig as ServerAuthenticator<Response<Empty<Bytes>>>>::get_server_layer(
164                    oidc,
165                )?;
166            Ok(AuthKind::Oidc(layer))
167        }
168        #[cfg(not(target_family = "windows"))]
169        ServerAuthConfig::Spire(spire) => {
170            let mut layer =
171                <SpireAuthConfig as ServerAuthenticator<Response<Empty<Bytes>>>>::get_server_layer(
172                    spire,
173                )?;
174            layer.initialize().await?;
175            Ok(AuthKind::Spire(layer))
176        }
177    }
178}
179
180impl ServerConfig {
181    pub async fn run_websocket_server(
182        &self,
183        drain_rx: drain::Watch,
184        on_accepted: OnAcceptedWebSocket,
185    ) -> Result<CancellationToken, ConfigError> {
186        if self.resolved_transport() != TransportProtocol::Websocket {
187            return Err(ConfigError::WebSocketServerUnsupportedTransport);
188        }
189
190        let endpoint = WebSocketEndpoint::parse(self.endpoint.as_str())?;
191        let listener = TcpListener::bind(endpoint.socket_address()).await?;
192
193        let tls_config = self.tls_setting.load_rustls_config().await?;
194        let tls_acceptor = match (endpoint.secure, tls_config) {
195            (true, Some(config)) => Some(TlsAcceptor::from(Arc::new(config))),
196            (true, None) => return Err(ConfigError::WebSocketServerTlsMissing),
197            (false, Some(_)) => return Err(ConfigError::WebSocketServerTlsUnexpected),
198            (false, None) => None,
199        };
200
201        let auth_kind = build_auth_kind(self).await?;
202        let expected_path = endpoint.path.clone();
203
204        // Cap the number of concurrent accept tasks
205        let max_connections = self
206            .max_concurrent_streams
207            .map(|n| n as usize)
208            .unwrap_or(DEFAULT_MAX_WEBSOCKET_CONNECTIONS)
209            .max(1);
210        let connection_semaphore = Arc::new(Semaphore::new(max_connections));
211
212        let cancellation_token = CancellationToken::new();
213        let cancel_clone = cancellation_token.clone();
214
215        tokio::spawn(async move {
216            let mut drain_signal = std::pin::pin!(drain_rx.signaled());
217
218            loop {
219                tokio::select! {
220                    _ = &mut drain_signal => {
221                        debug!("websocket server shutting down on drain");
222                        break;
223                    }
224                    _ = cancel_clone.cancelled() => {
225                        debug!("websocket server shutting down on cancellation token");
226                        break;
227                    }
228                    accepted = listener.accept() => {
229                        let (stream, remote_addr) = match accepted {
230                            Ok(val) => val,
231                            Err(err) => {
232                                warn!(error = %err, "websocket accept error");
233                                continue;
234                            }
235                        };
236
237                        // Acquire a slot before spawning. If the cap is
238                        // exhausted, drop the connection rather than
239                        // queueing it up
240                        let permit = match connection_semaphore.clone().try_acquire_owned() {
241                            Ok(permit) => permit,
242                            Err(_) => {
243                                warn!(
244                                    max_connections,
245                                    %remote_addr,
246                                    "websocket connection cap reached; rejecting client"
247                                );
248                                drop(stream);
249                                continue;
250                            }
251                        };
252
253                        let local_addr = stream.local_addr().ok();
254                        let auth_kind = auth_kind.clone();
255                        let expected_path = expected_path.clone();
256                        let on_accepted = on_accepted.clone();
257                        let tls_acceptor = tls_acceptor.clone();
258
259                        tokio::spawn(async move {
260                            let stream = match tls_acceptor {
261                                Some(acceptor) => {
262                                    // Bound TLS handshake duration so a
263                                    // silent/malicious client cannot hold a
264                                    // task forever.
265                                    match tokio::time::timeout(
266                                        TLS_HANDSHAKE_TIMEOUT,
267                                        acceptor.accept(stream),
268                                    )
269                                    .await
270                                    {
271                                        Ok(Ok(stream)) => MaybeTlsStream::Tls(Box::new(stream)),
272                                        Ok(Err(err)) => {
273                                            warn!(error = %err, "websocket TLS accept error");
274                                            // permit dropped here, slot released
275                                            return;
276                                        }
277                                        Err(_) => {
278                                            warn!(
279                                                timeout = ?TLS_HANDSHAKE_TIMEOUT,
280                                                "websocket TLS handshake timed out"
281                                            );
282                                            return;
283                                        }
284                                    }
285                                }
286                                None => MaybeTlsStream::Plain(stream),
287                            };
288
289                            // Ownership of `permit` is moved into
290                            // `serve_connection`, which in turn hands it
291                            // to the spawned websocket task so the slot
292                            // stays reserved for the lifetime of the
293                            // active websocket (not just the upgrade).
294                            serve_connection(
295                                stream,
296                                auth_kind,
297                                expected_path,
298                                on_accepted,
299                                remote_addr,
300                                local_addr,
301                                permit,
302                            )
303                            .await;
304                        });
305                    }
306                }
307            }
308        });
309
310        Ok(cancellation_token)
311    }
312}
313
314async fn serve_connection<S>(
315    stream: S,
316    auth_kind: AuthKind,
317    expected_path: String,
318    on_accepted: OnAcceptedWebSocket,
319    remote_addr: SocketAddr,
320    local_addr: Option<SocketAddr>,
321    permit: OwnedSemaphorePermit,
322) where
323    S: AsyncRead + AsyncWrite + Send + Unpin + 'static,
324{
325    let io = TokioIo::new(stream);
326
327    // Tracks whether the WebSocket upgrade succeeded. Used to enforce a
328    // timeout on the *upgrade* portion of the connection only; once the
329    // upgrade is complete the websocket itself may stay open indefinitely.
330    let upgrade_done = Arc::new(AtomicBool::new(false));
331    let upgrade_done_service = upgrade_done.clone();
332
333    // The connection permit is moved into the upgrade task on a successful
334    // upgrade so that the semaphore slot is held for the entire websocket
335    // lifetime. If the upgrade never happens (404/400/401/timeout), the
336    // permit is dropped here when `serve_connection` returns.
337    let permit_slot = Arc::new(parking_lot::Mutex::new(Some(permit)));
338    let permit_slot_service = permit_slot.clone();
339
340    let inner = service_fn(move |mut request: Request<Incoming>| {
341        let expected_path = expected_path.clone();
342        let on_accepted = on_accepted.clone();
343        let upgrade_done = upgrade_done_service.clone();
344        let permit_slot = permit_slot_service.clone();
345
346        let fut: Pin<Box<dyn Future<Output = Result<Response<Empty<Bytes>>, Infallible>> + Send>> =
347            Box::pin(async move {
348                if request.uri().path() != expected_path {
349                    return Ok::<Response<Empty<Bytes>>, Infallible>(response_with_status(
350                        StatusCode::NOT_FOUND,
351                    ));
352                }
353
354                if !upgrade::is_upgrade_request(&request) {
355                    return Ok::<Response<Empty<Bytes>>, Infallible>(response_with_status(
356                        StatusCode::BAD_REQUEST,
357                    ));
358                }
359
360                match upgrade::upgrade(&mut request) {
361                    Ok((response, future)) => {
362                        // Mark the upgrade as successful so the connection
363                        // future is no longer bound by `HTTP_UPGRADE_TIMEOUT`.
364                        upgrade_done.store(true, Ordering::SeqCst);
365                        // Transfer the semaphore permit out of the accept
366                        // task and into the websocket task: the connection
367                        // cap must bound *active websockets*, not just the
368                        // brief upgrade phase.
369                        let permit = permit_slot.lock().take();
370                        tokio::spawn(async move {
371                            let _permit = permit;
372                            match future.await {
373                                Ok(websocket) => {
374                                    on_accepted(AcceptedWebSocketConnection {
375                                        websocket,
376                                        remote_addr: Some(remote_addr),
377                                        local_addr,
378                                    })
379                                    .await;
380                                }
381                                Err(err) => {
382                                    warn!(error = %err, "websocket upgrade error");
383                                }
384                            }
385                        });
386
387                        Ok::<Response<Empty<Bytes>>, Infallible>(response)
388                    }
389                    Err(err) => {
390                        warn!(error = %err, "websocket upgrade rejected");
391                        Ok::<Response<Empty<Bytes>>, Infallible>(response_with_status(
392                            StatusCode::BAD_REQUEST,
393                        ))
394                    }
395                }
396            });
397        fut
398    });
399
400    // Bound the HTTP/WS upgrade phase so a client that opens a TCP/TLS
401    // connection but never sends a valid upgrade request cannot pin this
402    // task. The timeout only applies until `upgrade_done` flips: once the
403    // upgrade succeeds the websocket IO is hijacked into its own task and
404    // the remaining `serve_connection` future may resolve at any time
405    // without being killed by the upgrade timeout.
406    let upgrade_deadline = tokio::time::sleep(HTTP_UPGRADE_TIMEOUT);
407    tokio::pin!(upgrade_deadline);
408
409    // Each auth variant produces a different concrete tower `Service` type
410    // (the static Layer stacking precludes a single variable), so we
411    // type-erase each variant into a `BoxCloneService` with identical
412    // signatures and then hand it to hyper via `TowerToHyperService`. The
413    // dispatch mirrors the gRPC server's auth-layer construction
414    // (`grpc/server.rs`).
415    let svc: BoxCloneService<Request<Incoming>, Response<Empty<Bytes>>, Infallible> =
416        match auth_kind {
417            AuthKind::None => BoxCloneService::new(inner),
418            AuthKind::Basic(layer) => {
419                BoxCloneService::new(ServiceBuilder::new().layer(layer).service(inner))
420            }
421            AuthKind::Jwt(layer) => {
422                // QueryTokenToAuthHeaderLayer promotes `?token=<jwt>` to an
423                // `Authorization: Bearer <jwt>` header so browser clients —
424                // which cannot set custom headers on `new WebSocket(url)` —
425                // can still authenticate. Existing headers win over query.
426                BoxCloneService::new(
427                    ServiceBuilder::new()
428                        .layer(QueryTokenToAuthHeaderLayer::new())
429                        .layer(layer)
430                        .service(inner),
431                )
432            }
433            AuthKind::Oidc(layer) => BoxCloneService::new(
434                ServiceBuilder::new()
435                    .layer(QueryTokenToAuthHeaderLayer::new())
436                    .layer(layer)
437                    .service(inner),
438            ),
439            #[cfg(not(target_family = "windows"))]
440            AuthKind::Spire(layer) => BoxCloneService::new(
441                ServiceBuilder::new()
442                    .layer(QueryTokenToAuthHeaderLayer::new())
443                    .layer(layer)
444                    .service(inner),
445            ),
446        };
447
448    let connection = http1::Builder::new()
449        .serve_connection(io, TowerToHyperService::new(svc))
450        .with_upgrades();
451    tokio::pin!(connection);
452
453    tokio::select! {
454        biased;
455        result = &mut connection => {
456            if let Err(err) = result {
457                debug!(error = %err, "websocket HTTP connection closed with error");
458            }
459        }
460        _ = &mut upgrade_deadline, if !upgrade_done.load(Ordering::SeqCst) => {
461            warn!(
462                timeout = ?HTTP_UPGRADE_TIMEOUT,
463                "websocket HTTP upgrade timed out"
464            );
465        }
466    }
467}
468
469fn response_with_status(status: StatusCode) -> Response<Empty<Bytes>> {
470    Response::builder()
471        .status(status)
472        .body(Empty::new())
473        .expect("valid websocket HTTP response")
474}
475
476#[cfg(test)]
477mod tests {
478    use super::*;
479
480    use crate::auth::basic::Config as BasicConfig;
481    use crate::client::ClientConfig;
482    use crate::server::AuthenticationConfig as ServerAuthConfig;
483    use crate::tls::client::TlsClientConfig;
484    use crate::tls::server::TlsServerConfig;
485    use std::net::TcpListener as StdTcpListener;
486    use std::time::Duration;
487    use tokio::io::{AsyncReadExt, AsyncWriteExt};
488    use tokio::net::TcpStream as TokioTcpStream;
489
490    fn available_port() -> u16 {
491        StdTcpListener::bind("127.0.0.1:0")
492            .expect("bind")
493            .local_addr()
494            .expect("local_addr")
495            .port()
496    }
497
498    /// Poll for server readiness with exponential backoff (max ~2s) to avoid
499    /// flaky sleep-based readiness checks.
500    async fn wait_for_server_ready(addr: &str, max_attempts: u32) -> bool {
501        for attempt in 0..max_attempts {
502            if TokioTcpStream::connect(addr).await.is_ok() {
503                return true;
504            }
505            let backoff = Duration::from_millis(25 * (1 + attempt as u64).min(10));
506            tokio::time::sleep(backoff).await;
507        }
508        false
509    }
510
511    fn noop_on_accepted() -> OnAcceptedWebSocket {
512        Arc::new(|_| Box::pin(async {}))
513    }
514
515    async fn start_ws_server(server_conf: ServerConfig) -> CancellationToken {
516        let port = server_conf
517            .endpoint
518            .rsplit(':')
519            .next()
520            .and_then(|p| p.parse::<u16>().ok())
521            .expect("port");
522        let (signal, watch) = drain::channel();
523        // Keep the signal alive for the lifetime of the test; dropping it
524        // would immediately drain the server.
525        std::mem::forget(signal);
526        let token = server_conf
527            .run_websocket_server(watch, noop_on_accepted())
528            .await
529            .expect("server start");
530        assert!(
531            wait_for_server_ready(&format!("127.0.0.1:{port}"), 40).await,
532            "server did not become ready in time",
533        );
534        token
535    }
536
537    #[tokio::test]
538    async fn test_websocket_server_starts() {
539        let port = available_port();
540        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
541            .with_tls_settings(TlsServerConfig::insecure());
542
543        let token = start_ws_server(cfg).await;
544        token.cancel();
545    }
546
547    #[tokio::test]
548    async fn test_websocket_server_rejects_non_websocket_transport() {
549        let port = available_port();
550        // Default transport is gRPC.
551        let cfg = ServerConfig::with_endpoint(&format!("127.0.0.1:{port}"));
552        let (signal, watch) = drain::channel();
553        std::mem::forget(signal);
554        let res = cfg.run_websocket_server(watch, noop_on_accepted()).await;
555        assert!(matches!(
556            res,
557            Err(ConfigError::WebSocketServerUnsupportedTransport)
558        ));
559    }
560
561    #[tokio::test]
562    async fn test_websocket_server_rejects_invalid_endpoint() {
563        let cfg = ServerConfig::with_endpoint("not-a-ws-uri")
564            .with_tls_settings(TlsServerConfig::insecure());
565        let (signal, watch) = drain::channel();
566        std::mem::forget(signal);
567        let res = cfg.run_websocket_server(watch, noop_on_accepted()).await;
568        assert!(res.is_err());
569    }
570
571    async fn raw_http_request(addr: &str, request: &str) -> String {
572        let mut stream = TokioTcpStream::connect(addr).await.expect("tcp connect");
573        stream.write_all(request.as_bytes()).await.expect("write");
574        stream.flush().await.expect("flush");
575
576        let mut response = Vec::with_capacity(512);
577        let mut buf = [0u8; 256];
578        let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
579        loop {
580            let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
581            if remaining.is_zero() {
582                break;
583            }
584            match tokio::time::timeout(remaining, stream.read(&mut buf)).await {
585                Ok(Ok(0)) => break,
586                Ok(Ok(n)) => {
587                    response.extend_from_slice(&buf[..n]);
588                    if response.windows(4).any(|w| w == b"\r\n\r\n") {
589                        break;
590                    }
591                }
592                Ok(Err(_)) | Err(_) => break,
593            }
594        }
595        String::from_utf8_lossy(&response).to_string()
596    }
597
598    #[tokio::test]
599    async fn test_websocket_server_404_on_unknown_path() {
600        let port = available_port();
601        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
602            .with_tls_settings(TlsServerConfig::insecure());
603        let token = start_ws_server(cfg).await;
604
605        let req = format!(
606            "GET /wrong/path HTTP/1.1\r\n\
607             Host: 127.0.0.1:{port}\r\n\
608             Upgrade: websocket\r\n\
609             Connection: Upgrade\r\n\
610             Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
611             Sec-WebSocket-Version: 13\r\n\
612             \r\n",
613        );
614        let resp = raw_http_request(&format!("127.0.0.1:{port}"), &req).await;
615        assert!(
616            resp.starts_with("HTTP/1.1 404"),
617            "expected 404, got: {resp:?}"
618        );
619
620        token.cancel();
621    }
622
623    #[tokio::test]
624    async fn test_websocket_server_400_when_not_upgrade() {
625        let port = available_port();
626        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
627            .with_tls_settings(TlsServerConfig::insecure());
628        let token = start_ws_server(cfg).await;
629
630        // Correct path, but no Upgrade header.
631        let req = format!("GET / HTTP/1.1\r\nHost: 127.0.0.1:{port}\r\nConnection: close\r\n\r\n",);
632        let resp = raw_http_request(&format!("127.0.0.1:{port}"), &req).await;
633        assert!(
634            resp.starts_with("HTTP/1.1 400"),
635            "expected 400, got: {resp:?}"
636        );
637
638        token.cancel();
639    }
640
641    #[tokio::test]
642    async fn test_websocket_server_401_on_failed_basic_auth() {
643        let port = available_port();
644        let test_user = format!("user-{}", std::process::id());
645        let test_pass = format!("pw-{}-{}", std::process::id(), port);
646        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
647            .with_tls_settings(TlsServerConfig::insecure())
648            .with_auth(ServerAuthConfig::Basic(BasicConfig::new(
649                &test_user, &test_pass,
650            )));
651        let token = start_ws_server(cfg).await;
652
653        // Correct path & upgrade headers but no Authorization.
654        let req = format!(
655            "GET / HTTP/1.1\r\n\
656             Host: 127.0.0.1:{port}\r\n\
657             Upgrade: websocket\r\n\
658             Connection: Upgrade\r\n\
659             Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
660             Sec-WebSocket-Version: 13\r\n\
661             \r\n",
662        );
663        let resp = raw_http_request(&format!("127.0.0.1:{port}"), &req).await;
664        assert!(
665            resp.starts_with("HTTP/1.1 401"),
666            "expected 401, got: {resp:?}"
667        );
668
669        token.cancel();
670    }
671
672    /// Server must accept a WebSocket upgrade whose JWT is supplied via
673    /// `?token=<jwt>` query string instead of an `Authorization` header.
674    /// This is the browser-compatible path — `new WebSocket(url)` cannot
675    /// set custom headers, so the server's `QueryTokenToAuthHeaderLayer`
676    /// promotes the query param into a Bearer header before the JWT
677    /// validation layer runs.
678    #[tokio::test]
679    async fn test_websocket_server_accepts_jwt_via_query_token() {
680        use crate::auth::jwt::{Claims, Config as JwtConfig, JwtKey};
681        use slim_auth::jwt::{Algorithm, Key, KeyData, KeyFormat};
682        use slim_auth::traits::Signer;
683
684        let port = available_port();
685        let claims = Claims::new(
686            Some(vec!["audience".to_string()]),
687            Some("issuer".to_string()),
688            Some("subject".to_string()),
689            None,
690        );
691        let secret = format!("ws-jwt-secret-{}-{port}", std::process::id());
692        let encoding_key = JwtKey::Encoding(Key {
693            algorithm: Algorithm::HS256,
694            format: KeyFormat::Pem,
695            key: KeyData::Data(secret.clone()),
696        });
697        let decoding_key = JwtKey::Decoding(Key {
698            algorithm: Algorithm::HS256,
699            format: KeyFormat::Pem,
700            key: KeyData::Data(secret),
701        });
702        let client_jwt = JwtConfig::new(claims.clone(), Duration::from_secs(3600), encoding_key);
703        let server_jwt = JwtConfig::new(claims, Duration::from_secs(3600), decoding_key);
704
705        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
706            .with_tls_settings(TlsServerConfig::insecure())
707            .with_auth(ServerAuthConfig::Jwt(server_jwt));
708        let token = start_ws_server(cfg).await;
709
710        // Mint a JWT directly (simulates a browser that obtained one OOB).
711        let signer = client_jwt.get_provider().expect("signer");
712        let jwt = signer.sign_standard_claims().expect("sign");
713
714        // Send WS upgrade with `?token=<jwt>` and NO Authorization header.
715        let req = format!(
716            "GET /?token={jwt} HTTP/1.1\r\n\
717             Host: 127.0.0.1:{port}\r\n\
718             Upgrade: websocket\r\n\
719             Connection: Upgrade\r\n\
720             Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
721             Sec-WebSocket-Version: 13\r\n\
722             \r\n",
723        );
724        let resp = raw_http_request(&format!("127.0.0.1:{port}"), &req).await;
725        assert!(
726            resp.starts_with("HTTP/1.1 101"),
727            "expected 101 Switching Protocols, got: {resp:?}"
728        );
729
730        // Negative control: same upgrade without the query token must be
731        // rejected, confirming the 101 above is owed to the query param
732        // (not to the auth layer being absent).
733        let req_no_token = format!(
734            "GET / HTTP/1.1\r\n\
735             Host: 127.0.0.1:{port}\r\n\
736             Upgrade: websocket\r\n\
737             Connection: Upgrade\r\n\
738             Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
739             Sec-WebSocket-Version: 13\r\n\
740             \r\n",
741        );
742        let resp_no_token = raw_http_request(&format!("127.0.0.1:{port}"), &req_no_token).await;
743        assert!(
744            resp_no_token.starts_with("HTTP/1.1 401"),
745            "expected 401 without query token, got: {resp_no_token:?}"
746        );
747
748        token.cancel();
749    }
750
751    /// When both a Bearer Authorization header and a `?token=<jwt>` query
752    /// parameter are present, the header MUST win and the query MUST be
753    /// ignored. This pins the precedence so a client that already sets the
754    /// Authorization header cannot be silently overridden by a stale or
755    /// malicious URL query.
756    #[tokio::test]
757    async fn test_websocket_server_authorization_header_wins_over_query_token() {
758        use crate::auth::jwt::{Claims, Config as JwtConfig, JwtKey};
759        use slim_auth::jwt::{Algorithm, Key, KeyData, KeyFormat};
760        use slim_auth::traits::Signer;
761
762        let port = available_port();
763        let claims = Claims::new(
764            Some(vec!["audience".to_string()]),
765            Some("issuer".to_string()),
766            Some("subject".to_string()),
767            None,
768        );
769        let secret = format!("ws-jwt-secret-{}-{port}", std::process::id());
770        let encoding_key = JwtKey::Encoding(Key {
771            algorithm: Algorithm::HS256,
772            format: KeyFormat::Pem,
773            key: KeyData::Data(secret.clone()),
774        });
775        let decoding_key = JwtKey::Decoding(Key {
776            algorithm: Algorithm::HS256,
777            format: KeyFormat::Pem,
778            key: KeyData::Data(secret),
779        });
780        let client_jwt = JwtConfig::new(claims.clone(), Duration::from_secs(3600), encoding_key);
781        let server_jwt = JwtConfig::new(claims, Duration::from_secs(3600), decoding_key);
782
783        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
784            .with_tls_settings(TlsServerConfig::insecure())
785            .with_auth(ServerAuthConfig::Jwt(server_jwt));
786        let token = start_ws_server(cfg).await;
787
788        let signer = client_jwt.get_provider().expect("signer");
789        let valid_jwt = signer.sign_standard_claims().expect("sign");
790
791        // Bogus Authorization header + valid `?token=<jwt>`. If the query
792        // were used the upgrade would succeed (101). Because the header
793        // wins, JWT validation rejects the garbage Bearer and returns 401.
794        let req = format!(
795            "GET /?token={valid_jwt} HTTP/1.1\r\n\
796             Host: 127.0.0.1:{port}\r\n\
797             Authorization: Bearer not-a-valid-jwt\r\n\
798             Upgrade: websocket\r\n\
799             Connection: Upgrade\r\n\
800             Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
801             Sec-WebSocket-Version: 13\r\n\
802             \r\n",
803        );
804        let resp = raw_http_request(&format!("127.0.0.1:{port}"), &req).await;
805        assert!(
806            resp.starts_with("HTTP/1.1 401"),
807            "header must win over query token (expected 401), got: {resp:?}"
808        );
809
810        // Symmetric positive control: valid Bearer + bogus query → 101.
811        // Confirms the rejection above is due to the bad header, not the
812        // presence of the query param itself.
813        let req_valid_header = format!(
814            "GET /?token=not-a-valid-jwt HTTP/1.1\r\n\
815             Host: 127.0.0.1:{port}\r\n\
816             Authorization: Bearer {valid_jwt}\r\n\
817             Upgrade: websocket\r\n\
818             Connection: Upgrade\r\n\
819             Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
820             Sec-WebSocket-Version: 13\r\n\
821             \r\n",
822        );
823        let resp_valid_header =
824            raw_http_request(&format!("127.0.0.1:{port}"), &req_valid_header).await;
825        assert!(
826            resp_valid_header.starts_with("HTTP/1.1 101"),
827            "valid header must succeed regardless of query (expected 101), got: {resp_valid_header:?}"
828        );
829
830        token.cancel();
831    }
832
833    #[tokio::test]
834    async fn test_websocket_server_full_handshake_via_client() {
835        let port = available_port();
836        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
837            .with_tls_settings(TlsServerConfig::insecure());
838        let token = start_ws_server(cfg).await;
839
840        let client_cfg = ClientConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
841            .with_tls_setting(TlsClientConfig::insecure());
842
843        let channel =
844            tokio::time::timeout(Duration::from_secs(5), client_cfg.to_websocket_channel())
845                .await
846                .expect("handshake timed out")
847                .expect("handshake failed");
848
849        assert!(channel.remote_addr().is_some());
850        // local_addr is unavailable when driving the upgrade via hyper-util's
851        // legacy Client — its HttpInfo extension only exposes remote_addr.
852        assert!(channel.local_addr().is_none());
853
854        token.cancel();
855    }
856
857    #[tokio::test]
858    async fn test_websocket_server_cancellation_stops_listener() {
859        let port = available_port();
860        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
861            .with_tls_settings(TlsServerConfig::insecure());
862        let token = start_ws_server(cfg).await;
863
864        token.cancel();
865
866        // Wait for listener to actually stop accepting.
867        for _ in 0..40 {
868            if TokioTcpStream::connect(format!("127.0.0.1:{port}"))
869                .await
870                .is_err()
871            {
872                return;
873            }
874            tokio::time::sleep(Duration::from_millis(50)).await;
875        }
876        panic!("listener did not stop after cancellation");
877    }
878
879    #[tokio::test]
880    async fn test_websocket_server_enforces_max_connections() {
881        let port = available_port();
882        let cfg = ServerConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
883            .with_tls_settings(TlsServerConfig::insecure())
884            .with_max_concurrent_streams(Some(1));
885
886        let on_accepted: OnAcceptedWebSocket = Arc::new(|_| {
887            Box::pin(async {
888                tokio::time::sleep(Duration::from_secs(30)).await;
889            })
890        });
891
892        let (signal, watch) = drain::channel();
893        std::mem::forget(signal);
894        let token = cfg
895            .run_websocket_server(watch, on_accepted)
896            .await
897            .expect("server start");
898
899        let client_cfg = ClientConfig::with_endpoint(&format!("ws://127.0.0.1:{port}"))
900            .with_tls_setting(TlsClientConfig::insecure())
901            .with_connect_timeout(Duration::from_secs(5))
902            .with_backoff(crate::client::BackoffConfig::new_fixed_interval(
903                Duration::from_millis(0),
904                1,
905            ));
906
907        let mut first = None;
908        for _ in 0..40 {
909            match tokio::time::timeout(Duration::from_secs(5), client_cfg.to_websocket_channel())
910                .await
911            {
912                Ok(Ok(ch)) => {
913                    first = Some(ch);
914                    break;
915                }
916                _ => tokio::time::sleep(Duration::from_millis(50)).await,
917            }
918        }
919        let _hold_first = first.expect("first handshake never succeeded");
920
921        tokio::time::sleep(Duration::from_millis(200)).await;
922
923        let second =
924            tokio::time::timeout(Duration::from_secs(3), client_cfg.to_websocket_channel()).await;
925        let inner = second.expect("second handshake outer timeout");
926        assert!(
927            inner.is_err(),
928            "second handshake must fail when max_concurrent_streams cap is reached"
929        );
930
931        token.cancel();
932    }
933}