1use 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
55const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
59
60const HTTP_UPGRADE_TIMEOUT: Duration = Duration::from_secs(10);
64
65const DEFAULT_MAX_WEBSOCKET_CONNECTIONS: usize = 1024;
70
71enum 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 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#[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 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 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 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 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 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 let upgrade_done = Arc::new(AtomicBool::new(false));
331 let upgrade_done_service = upgrade_done.clone();
332
333 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 upgrade_done.store(true, Ordering::SeqCst);
365 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 let upgrade_deadline = tokio::time::sleep(HTTP_UPGRADE_TIMEOUT);
407 tokio::pin!(upgrade_deadline);
408
409 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 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 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 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 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 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 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 #[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 let signer = client_jwt.get_provider().expect("signer");
712 let jwt = signer.sign_standard_claims().expect("sign");
713
714 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 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 #[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 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 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 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 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}