1#[cfg(feature = "relay-client")]
16pub mod client;
17pub mod protocol;
18pub mod proxy;
19pub mod registry;
20
21use std::net::SocketAddr;
22use std::sync::Arc;
23use std::time::Duration;
24
25use axum::{
26 body::Bytes,
27 extract::{
28 ws::{Message, WebSocket, WebSocketUpgrade},
29 FromRequestParts, Request, State,
30 },
31 http::{HeaderMap, StatusCode},
32 response::{IntoResponse, Response},
33 routing::{any, get},
34 Router,
35};
36use futures_util::{SinkExt, StreamExt};
37
38use crate::error::ShellTunnelError;
39use crate::security::{generate_api_key, rate_limit_middleware, RateLimitConfig, RateLimiter};
40use protocol::{reject, DeviceMessage, RelayMessage, PROTOCOL_VERSION};
41use proxy::{
42 is_forwardable, split_device_path, ProxyRequest, ProxyResponse, POOL_WAIT, REQUEST_TIMEOUT,
43};
44use registry::{Device, DeviceRegistry};
45
46pub use registry::{DeviceRegistry as Registry, POOL_TARGET};
47
48pub const HEARTBEAT_TIMEOUT: Duration = Duration::from_secs(90);
50
51const ENROLL_TIMEOUT: Duration = Duration::from_secs(10);
53
54#[derive(Debug, Clone)]
56pub struct RelayConfig {
57 pub bind: SocketAddr,
59 pub enroll_token: String,
61 pub rate_limit: RateLimitConfig,
68 #[cfg(feature = "tls")]
70 pub tls: Option<crate::tls::TlsFiles>,
71 pub public_base: Option<String>,
77}
78
79impl RelayConfig {
80 pub fn new(bind: SocketAddr, enroll_token: impl Into<String>) -> Self {
82 Self {
83 bind,
84 enroll_token: enroll_token.into(),
85 rate_limit: RateLimitConfig::default(),
86 #[cfg(feature = "tls")]
87 tls: None,
88 public_base: None,
89 }
90 }
91
92 #[cfg(feature = "tls")]
94 pub fn with_tls(mut self, files: crate::tls::TlsFiles) -> Self {
95 self.tls = Some(files);
96 self
97 }
98
99 pub fn without_rate_limit(mut self) -> Self {
101 self.rate_limit.enabled = false;
102 self
103 }
104
105 pub fn with_public_base(mut self, base: impl Into<String>) -> Self {
107 self.public_base = Some(base.into().trim_end_matches('/').to_string());
108 self
109 }
110
111 pub fn resolved_public_base(&self) -> Option<String> {
122 self.public_base.as_deref().map(|base| {
123 public_base_port_hint(base, self.bind.port()).unwrap_or_else(|| base.to_string())
124 })
125 }
126
127 pub fn public_base_or(&self, observed: Option<String>) -> String {
133 self.resolved_public_base()
134 .or(observed)
135 .unwrap_or_else(|| format!("http://{}", self.bind))
136 }
137
138 pub fn public_url_for(&self, device_id: &str, observed: Option<String>) -> String {
140 format!("{}/d/{}", self.public_base_or(observed), device_id)
141 }
142}
143
144pub fn public_base_port_hint(base: &str, listen_port: u16) -> Option<String> {
154 let (scheme, rest) = base.split_once("://")?;
155 let default_port: u16 = match scheme {
156 "https" => 443,
157 "http" => 80,
158 _ => return None,
159 };
160 if listen_port == default_port {
161 return None;
162 }
163 let authority_end = rest.find('/').unwrap_or(rest.len());
164 let authority = &rest[..authority_end];
165 let has_port = match authority.rfind(']') {
168 Some(bracket) => authority[bracket..].contains(':'),
169 None => authority.contains(':'),
170 };
171 if has_port {
172 return None;
173 }
174 Some(format!(
175 "{scheme}://{authority}:{listen_port}{}",
176 &rest[authority_end..]
177 ))
178}
179
180#[derive(Debug, Clone)]
182pub struct RelayState {
183 config: RelayConfig,
184 devices: DeviceRegistry,
185}
186
187impl RelayState {
188 pub fn new(config: RelayConfig) -> Self {
190 Self {
191 config,
192 devices: DeviceRegistry::new(),
193 }
194 }
195
196 pub fn devices(&self) -> &DeviceRegistry {
198 &self.devices
199 }
200}
201
202pub fn relay_router(state: RelayState) -> Router {
208 let limiter = Arc::new(RateLimiter::new(state.config.rate_limit.clone()));
209
210 Router::new()
211 .route("/health", get(|| async { "OK" }))
212 .route("/relay/v1/control", get(control_handler))
213 .route("/relay/v1/data", get(data_handler))
214 .route("/relay/v1/devices", get(devices_handler))
215 .route("/d/{*rest}", any(proxy_handler))
216 .layer(axum::middleware::from_fn_with_state(
217 limiter,
218 rate_limit_middleware,
219 ))
220 .with_state(state)
221}
222
223pub async fn bind_relay(config: &RelayConfig) -> crate::Result<tokio::net::TcpListener> {
232 tokio::net::TcpListener::bind(config.bind)
233 .await
234 .map_err(crate::error::ShellTunnelError::Io)
235}
236
237pub async fn serve_relay(config: RelayConfig) -> crate::Result<()> {
243 let listener = bind_relay(&config).await?;
244 serve_relay_on(listener, config).await
245}
246
247pub async fn serve_relay_on(
249 listener: tokio::net::TcpListener,
250 config: RelayConfig,
251) -> crate::Result<()> {
252 let bind = config.bind;
253 #[cfg(feature = "tls")]
254 let tls = config.tls.clone();
255 let state = RelayState::new(config);
256 let router = relay_router(state.clone());
257
258 let sweeper = state.devices().clone();
261 tokio::spawn(async move {
262 let mut ticker = tokio::time::interval(HEARTBEAT_TIMEOUT / 3);
263 loop {
264 ticker.tick().await;
265 for id in sweeper.evict_stale(HEARTBEAT_TIMEOUT) {
266 tracing::info!(target: "relay", device_id = %id, "device evicted (no heartbeat)");
267 }
268 }
269 });
270
271 tracing::info!("relay listening on {}", bind);
274
275 let service = router.into_make_service_with_connect_info::<SocketAddr>();
278
279 #[cfg(feature = "tls")]
280 if let Some(files) = tls {
281 let config = crate::tls::acceptor(files.load()?);
284 crate::tls::watch(files, config.clone());
286 let std_listener = listener.into_std().map_err(ShellTunnelError::Io)?;
287 return axum_server::from_tcp_rustls(std_listener, config)
288 .map_err(ShellTunnelError::Io)?
289 .serve(service)
290 .await
291 .map_err(|e| ShellTunnelError::Io(std::io::Error::other(e.to_string())));
292 }
293
294 axum::serve(listener, service)
295 .await
296 .map_err(|e| ShellTunnelError::Io(std::io::Error::other(e.to_string())))?;
297 Ok(())
298}
299
300fn is_valid_device_name(name: &str) -> bool {
306 !name.is_empty()
307 && name.len() <= 64
308 && name
309 .chars()
310 .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_')
311}
312
313fn observed_base(headers: &HeaderMap, tls: bool) -> Option<String> {
319 let host = headers
320 .get("x-forwarded-host")
321 .or_else(|| headers.get(axum::http::header::HOST))
322 .and_then(|value| value.to_str().ok())?;
323 if host.is_empty() {
324 return None;
325 }
326 let scheme = headers
330 .get("x-forwarded-proto")
331 .and_then(|value| value.to_str().ok())
332 .map(|proto| proto.split(',').next().unwrap_or(proto).trim().to_string())
333 .unwrap_or_else(|| if tls { "https" } else { "http" }.to_string());
334 Some(format!("{scheme}://{host}"))
335}
336
337fn serves_tls(_state: &RelayState) -> bool {
339 #[cfg(feature = "tls")]
340 {
341 _state.config.tls.is_some()
342 }
343 #[cfg(not(feature = "tls"))]
344 {
345 false
346 }
347}
348
349async fn control_handler(
351 ws: WebSocketUpgrade,
352 State(state): State<RelayState>,
353 headers: HeaderMap,
354) -> impl IntoResponse {
355 let observed = observed_base(&headers, serves_tls(&state));
356 ws.on_upgrade(move |socket| control_session(socket, state, observed))
357}
358
359async fn control_session(socket: WebSocket, state: RelayState, observed: Option<String>) {
361 let (mut sink, mut stream) = socket.split();
362
363 let first = match tokio::time::timeout(ENROLL_TIMEOUT, stream.next()).await {
366 Ok(Some(Ok(Message::Text(text)))) => text,
367 _ => return,
368 };
369
370 let enroll = match serde_json::from_str::<DeviceMessage>(&first) {
371 Ok(DeviceMessage::Enroll {
372 enroll_token,
373 version,
374 label,
375 device_name,
376 }) => (enroll_token, version, label, device_name),
377 _ => {
378 reject_and_close(
379 &mut sink,
380 reject::BAD_HANDSHAKE,
381 "expected an enroll message",
382 )
383 .await;
384 return;
385 }
386 };
387 let (enroll_token, version, label, device_name) = enroll;
388
389 if version != PROTOCOL_VERSION {
390 reject_and_close(
391 &mut sink,
392 reject::UNSUPPORTED_VERSION,
393 &format!("relay speaks protocol version {PROTOCOL_VERSION}"),
394 )
395 .await;
396 return;
397 }
398
399 if !constant_time_eq(&enroll_token, &state.config.enroll_token) {
400 tracing::debug!(target: "relay", "enrollment rejected: bad token");
402 reject_and_close(&mut sink, reject::BAD_TOKEN, "enrollment refused").await;
403 return;
404 }
405
406 let device_id = match device_name {
411 Some(name) if !is_valid_device_name(&name) => {
412 reject_and_close(
413 &mut sink,
414 reject::BAD_DEVICE_NAME,
415 "device names may use letters, digits, '-' and '_' (1-64 characters)",
416 )
417 .await;
418 return;
419 }
420 Some(name) => name,
426 None => generate_api_key(),
427 };
428 let public_url = state.config.public_url_for(&device_id, observed);
429 let registry::DeviceHandles {
430 device,
431 mut refill_rx,
432 } = state.devices.attach(&device_id, label.clone());
433 tracing::info!(
434 target: "relay",
435 device_id = %device_id,
436 label = label.as_deref().unwrap_or("-"),
437 "device attached"
438 );
439
440 let enrolled = RelayMessage::Enrolled {
441 device_id: device_id.clone(),
442 public_url,
443 };
444 if send_json(&mut sink, &enrolled).await.is_err() {
445 state.devices.detach(&device_id);
446 return;
447 }
448
449 let fill = RelayMessage::OpenData {
451 count: registry::POOL_TARGET,
452 };
453 if send_json(&mut sink, &fill).await.is_err() {
454 state.devices.detach(&device_id);
455 return;
456 }
457
458 loop {
461 tokio::select! {
462 incoming = stream.next() => {
463 let Some(Ok(message)) = incoming else { break };
464 match message {
465 Message::Text(text) => match serde_json::from_str::<DeviceMessage>(&text) {
466 Ok(DeviceMessage::Heartbeat) => {
467 device.touch();
468 if send_json(&mut sink, &RelayMessage::HeartbeatAck).await.is_err() {
469 break;
470 }
471 }
472 _ => continue,
476 },
477 Message::Close(_) => break,
478 _ => continue,
479 }
480 }
481 refill = refill_rx.recv() => {
482 if refill.is_none() {
483 break;
484 }
485 if send_json(&mut sink, &RelayMessage::OpenData { count: 1 }).await.is_err() {
486 break;
487 }
488 }
489 }
490 }
491
492 state.devices.detach(&device_id);
493 tracing::info!(target: "relay", device_id = %device_id, "device detached");
494}
495
496async fn devices_handler(State(state): State<RelayState>, headers: HeaderMap) -> Response {
502 let presented = headers
503 .get(axum::http::header::AUTHORIZATION)
504 .and_then(|value| value.to_str().ok())
505 .and_then(|value| value.strip_prefix("Bearer "))
506 .unwrap_or("");
507 if !constant_time_eq(presented, &state.config.enroll_token) {
508 return StatusCode::UNAUTHORIZED.into_response();
509 }
510
511 let base = state
512 .config
513 .public_base_or(observed_base(&headers, serves_tls(&state)));
514 let devices: Vec<_> = state
515 .devices
516 .list()
517 .into_iter()
518 .map(|device| {
519 let url = format!("{}/d/{}", base, device.id);
520 serde_json::json!({
521 "id": device.id,
522 "label": device.label,
523 "attached_secs": device.attached_secs,
524 "last_seen_secs": device.last_seen_secs,
525 "public_url": url,
526 })
527 })
528 .collect();
529
530 axum::Json(serde_json::json!({ "devices": devices })).into_response()
531}
532
533async fn data_handler(ws: WebSocketUpgrade, State(state): State<RelayState>) -> Response {
540 ws.on_upgrade(move |socket| attach_data_connection(socket, state))
541}
542
543async fn attach_data_connection(mut socket: WebSocket, state: RelayState) {
545 let first = tokio::time::timeout(ENROLL_TIMEOUT, socket.recv()).await;
546 let Ok(Some(Ok(Message::Text(text)))) = first else {
547 let _ = socket.close().await;
548 return;
549 };
550
551 let Ok(DeviceMessage::Attach {
552 device_id,
553 enroll_token,
554 }) = serde_json::from_str::<DeviceMessage>(&text)
555 else {
556 let _ = socket.close().await;
557 return;
558 };
559
560 if !constant_time_eq(&enroll_token, &state.config.enroll_token) {
561 tracing::debug!(target: "relay", "data connection rejected: bad token");
562 let _ = socket.close().await;
563 return;
564 }
565
566 let Some(device) = state.devices.get(&device_id) else {
567 let _ = socket.close().await;
568 return;
569 };
570
571 if let Some(mut extra) = device.offer(socket).await {
574 let _ = extra.close().await;
575 }
576}
577
578async fn proxy_handler(State(state): State<RelayState>, request: Request) -> Response {
580 let path_and_query = request
581 .uri()
582 .path_and_query()
583 .map(|p| p.as_str().to_string())
584 .unwrap_or_else(|| request.uri().path().to_string());
585
586 let Some((device_id, tail)) = split_device_path(&path_and_query) else {
587 return StatusCode::NOT_FOUND.into_response();
588 };
589
590 let Some(device) = state.devices.get(device_id) else {
591 return (StatusCode::BAD_GATEWAY, "device is not connected").into_response();
594 };
595
596 let method = request.method().to_string();
597 let headers: Vec<(String, String)> = request
598 .headers()
599 .iter()
600 .filter(|(name, _)| is_forwardable(name.as_str()))
601 .filter_map(|(name, value)| {
602 value
603 .to_str()
604 .ok()
605 .map(|v| (name.as_str().to_string(), v.to_string()))
606 })
607 .collect();
608
609 if is_websocket_upgrade(request.headers()) {
615 let (mut parts, _) = request.into_parts();
616 let upgrade = match WebSocketUpgrade::from_request_parts(&mut parts, &state).await {
617 Ok(upgrade) => upgrade,
618 Err(rejection) => return rejection.into_response(),
619 };
620 let proxied = ProxyRequest {
621 method,
622 path: tail,
623 headers,
624 websocket: true,
625 };
626 return upgrade.on_upgrade(move |client| pipe_websocket(client, device, proxied));
627 }
628
629 let body = match axum::body::to_bytes(request.into_body(), MAX_BODY).await {
630 Ok(body) => body,
631 Err(_) => return StatusCode::PAYLOAD_TOO_LARGE.into_response(),
632 };
633
634 let Some(conn) = device.take(POOL_WAIT).await else {
635 return (
638 StatusCode::SERVICE_UNAVAILABLE,
639 [("retry-after", "1")],
640 "no data connection available",
641 )
642 .into_response();
643 };
644
645 match tokio::time::timeout(
646 REQUEST_TIMEOUT,
647 forward(
648 conn,
649 ProxyRequest {
650 method,
651 path: tail,
652 headers,
653 websocket: false,
654 },
655 body,
656 ),
657 )
658 .await
659 {
660 Ok(Ok(response)) => response,
661 Ok(Err(reason)) => {
662 tracing::debug!(target: "relay", device_id = %device.id, reason, "proxy failed");
663 (StatusCode::BAD_GATEWAY, "device did not answer").into_response()
664 }
665 Err(_) => (StatusCode::GATEWAY_TIMEOUT, "device timed out").into_response(),
666 }
667}
668
669fn is_websocket_upgrade(headers: &HeaderMap) -> bool {
671 let header_contains = |name: axum::http::HeaderName, needle: &str| {
672 headers
673 .get(name)
674 .and_then(|value| value.to_str().ok())
675 .is_some_and(|value| value.to_ascii_lowercase().contains(needle))
676 };
677 header_contains(axum::http::header::UPGRADE, "websocket")
678 && header_contains(axum::http::header::CONNECTION, "upgrade")
679}
680
681async fn pipe_websocket(mut client: WebSocket, device: Arc<Device>, request: ProxyRequest) {
687 let Some(mut conn) = device.take(POOL_WAIT).await else {
688 tracing::debug!(target: "relay", device_id = %device.id, "no data connection for websocket");
689 let _ = client.close().await;
690 return;
691 };
692
693 let Ok(header) = serde_json::to_string(&request) else {
694 let _ = client.close().await;
695 return;
696 };
697 if conn.send(Message::Text(header.into())).await.is_err() {
698 let _ = client.close().await;
699 return;
700 }
701
702 let switched = matches!(
705 conn.recv().await,
706 Some(Ok(Message::Text(ref text)))
707 if serde_json::from_str::<ProxyResponse>(text)
708 .map(|response| response.status == 101)
709 .unwrap_or(false)
710 );
711 if !switched {
712 let _ = client.close().await;
713 let _ = conn.close().await;
714 return;
715 }
716
717 loop {
720 tokio::select! {
721 from_client = client.recv() => {
722 match from_client {
723 Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
724 Some(Ok(message)) => {
725 if conn.send(message).await.is_err() {
726 break;
727 }
728 }
729 }
730 }
731 from_device = conn.recv() => {
732 match from_device {
733 Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
734 Some(Ok(message)) => {
735 if client.send(message).await.is_err() {
736 break;
737 }
738 }
739 }
740 }
741 }
742 }
743
744 let _ = client.close().await;
745 let _ = conn.close().await;
746}
747
748const MAX_BODY: usize = 8 * 1024 * 1024;
750
751async fn forward(
756 mut conn: WebSocket,
757 request: ProxyRequest,
758 body: Bytes,
759) -> Result<Response, &'static str> {
760 let header = serde_json::to_string(&request).map_err(|_| "request-encode")?;
761 conn.send(Message::Text(header.into()))
762 .await
763 .map_err(|_| "request-header-send")?;
764 conn.send(Message::Binary(body))
765 .await
766 .map_err(|_| "request-body-send")?;
767
768 let head: ProxyResponse = loop {
769 match conn.recv().await {
770 Some(Ok(Message::Text(text))) => {
771 break serde_json::from_str(&text).map_err(|_| "response-decode")?
772 }
773 Some(Ok(_)) => continue,
774 _ => return Err("response-header-missing"),
775 }
776 };
777
778 let mut body = Vec::new();
787 loop {
788 match conn.recv().await {
789 Some(Ok(Message::Binary(chunk))) => body.extend_from_slice(&chunk),
790 Some(Ok(Message::Close(_))) | None => break,
791 Some(Ok(_)) => continue,
792 Some(Err(_)) => return Err("response-body-truncated"),
793 }
794 }
795
796 let mut response = Response::builder().status(head.status);
797 for (name, value) in head.headers {
798 if is_forwardable(&name) {
799 response = response.header(name, value);
800 }
801 }
802 response
803 .body(axum::body::Body::from(body))
804 .map_err(|_| "response-build")
805}
806
807async fn reject_and_close<S>(sink: &mut S, code: &str, message: &str)
809where
810 S: SinkExt<Message> + Unpin,
811{
812 let rejected = RelayMessage::Rejected {
813 code: code.to_string(),
814 message: message.to_string(),
815 };
816 let _ = send_json(sink, &rejected).await;
817 let _ = sink.close().await;
818}
819
820async fn send_json<S, T>(sink: &mut S, message: &T) -> Result<(), ()>
822where
823 S: SinkExt<Message> + Unpin,
824 T: serde::Serialize,
825{
826 let json = serde_json::to_string(message).map_err(|_| ())?;
827 sink.send(Message::Text(json.into())).await.map_err(|_| ())
828}
829
830fn constant_time_eq(a: &str, b: &str) -> bool {
836 let (a, b) = (a.as_bytes(), b.as_bytes());
837 if a.len() != b.len() {
838 return false;
839 }
840 a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
841}
842
843#[cfg(test)]
844mod tests {
845 use super::*;
846
847 fn config() -> RelayConfig {
848 RelayConfig::new("127.0.0.1:0".parse().unwrap(), "secret")
849 }
850
851 #[test]
852 fn a_portless_base_on_a_nondefault_port_gets_a_corrected_suggestion() {
853 assert_eq!(
859 public_base_port_hint("https://labs.example.com", 8443).as_deref(),
860 Some("https://labs.example.com:8443")
861 );
862 assert_eq!(
863 public_base_port_hint("http://relay.local", 8080).as_deref(),
864 Some("http://relay.local:8080")
865 );
866 }
867
868 #[test]
869 fn a_base_matching_the_scheme_default_needs_no_hint() {
870 assert_eq!(public_base_port_hint("https://labs.example.com", 443), None);
871 assert_eq!(public_base_port_hint("http://relay.local", 80), None);
872 }
873
874 #[test]
875 fn an_explicit_port_is_the_operator_stating_intent() {
876 assert_eq!(
878 public_base_port_hint("https://labs.example.com:8443", 8443),
879 None
880 );
881 assert_eq!(
882 public_base_port_hint("https://labs.example.com:9000", 8443),
883 None
884 );
885 assert_eq!(
886 public_base_port_hint("https://labs.example.com:443", 8443),
887 None
888 );
889 }
890
891 #[test]
892 fn the_port_is_spliced_into_the_authority_not_the_tail() {
893 assert_eq!(
895 public_base_port_hint("https://labs.example.com/relay", 8443).as_deref(),
896 Some("https://labs.example.com:8443/relay")
897 );
898 }
899
900 #[test]
901 fn ipv6_literals_look_for_the_port_after_the_bracket() {
902 assert_eq!(
903 public_base_port_hint("https://[::1]", 8443).as_deref(),
904 Some("https://[::1]:8443")
905 );
906 assert_eq!(public_base_port_hint("https://[::1]:8443", 8443), None);
907 }
908
909 #[test]
910 fn an_unrecognized_scheme_is_left_alone() {
911 assert_eq!(public_base_port_hint("ws://relay.local", 8443), None);
912 }
913
914 #[test]
915 fn public_url_uses_the_device_path_prefix() {
916 let config = RelayConfig::new("127.0.0.1:443".parse().unwrap(), "secret")
919 .with_public_base("https://relay.example.com/");
920 assert_eq!(
921 config.public_url_for("dev-1", None),
922 "https://relay.example.com/d/dev-1"
923 );
924 }
925
926 #[test]
927 fn a_portless_base_inherits_the_listen_port() {
928 let config = RelayConfig::new("0.0.0.0:8443".parse().unwrap(), "secret")
932 .with_public_base("https://labs.example.com");
933 assert_eq!(
934 config.resolved_public_base().as_deref(),
935 Some("https://labs.example.com:8443")
936 );
937 assert_eq!(
938 config.public_url_for("dev-1", None),
939 "https://labs.example.com:8443/d/dev-1"
940 );
941 }
942
943 #[test]
944 fn an_explicit_port_survives_resolution() {
945 let config = RelayConfig::new("0.0.0.0:8443".parse().unwrap(), "secret")
948 .with_public_base("https://labs.example.com:443");
949 assert_eq!(
950 config.resolved_public_base().as_deref(),
951 Some("https://labs.example.com:443")
952 );
953 }
954
955 #[test]
956 fn resolution_leaves_a_default_port_base_alone() {
957 let config = RelayConfig::new("0.0.0.0:443".parse().unwrap(), "secret")
959 .with_public_base("https://labs.example.com");
960 assert_eq!(
961 config.resolved_public_base().as_deref(),
962 Some("https://labs.example.com")
963 );
964 }
965
966 #[test]
967 fn public_base_defaults_to_the_bind_address() {
968 let config = RelayConfig::new("127.0.0.1:8443".parse().unwrap(), "secret");
969 assert_eq!(
970 config.public_url_for("d", None),
971 "http://127.0.0.1:8443/d/d"
972 );
973 }
974
975 #[test]
976 fn an_observed_address_is_used_when_the_operator_configured_none() {
977 let config = config();
978 assert_eq!(
979 config.public_url_for("dev-1", Some("https://relay.example.com".into())),
980 "https://relay.example.com/d/dev-1"
981 );
982 }
983
984 #[test]
985 fn a_configured_base_wins_over_what_the_connection_observed() {
986 let config = RelayConfig::new("127.0.0.1:443".parse().unwrap(), "secret")
987 .with_public_base("https://canonical.example");
988 assert_eq!(
989 config.public_url_for("dev-1", Some("https://whatever.invalid".into())),
990 "https://canonical.example/d/dev-1"
991 );
992 }
993
994 #[test]
995 fn the_forwarded_scheme_and_host_are_preferred_over_the_direct_host() {
996 let mut headers = HeaderMap::new();
997 headers.insert(axum::http::header::HOST, "127.0.0.1:8443".parse().unwrap());
998 assert_eq!(
999 observed_base(&headers, false).as_deref(),
1000 Some("http://127.0.0.1:8443")
1001 );
1002
1003 headers.insert("x-forwarded-proto", "https".parse().unwrap());
1004 headers.insert("x-forwarded-host", "relay.example.com".parse().unwrap());
1005 assert_eq!(
1006 observed_base(&headers, false).as_deref(),
1007 Some("https://relay.example.com")
1008 );
1009 }
1010
1011 #[test]
1012 fn a_proxy_chain_scheme_takes_the_first_entry() {
1013 let mut headers = HeaderMap::new();
1014 headers.insert(
1015 axum::http::header::HOST,
1016 "relay.example.com".parse().unwrap(),
1017 );
1018 headers.insert("x-forwarded-proto", "https, http".parse().unwrap());
1019 assert_eq!(
1020 observed_base(&headers, false).as_deref(),
1021 Some("https://relay.example.com")
1022 );
1023 }
1024
1025 #[test]
1026 fn no_host_header_means_nothing_observed() {
1027 assert!(observed_base(&HeaderMap::new(), false).is_none());
1028 }
1029
1030 #[test]
1031 fn terminating_tls_makes_the_advertised_url_https() {
1032 let mut headers = HeaderMap::new();
1035 headers.insert(
1036 axum::http::header::HOST,
1037 "relay.example.com".parse().unwrap(),
1038 );
1039 assert_eq!(
1040 observed_base(&headers, true).as_deref(),
1041 Some("https://relay.example.com")
1042 );
1043 }
1044
1045 #[test]
1046 fn device_names_must_be_url_path_safe() {
1047 assert!(is_valid_device_name("build-box"));
1048 assert!(is_valid_device_name("laptop_2"));
1049 assert!(is_valid_device_name("a"));
1050
1051 assert!(!is_valid_device_name(""));
1052 assert!(!is_valid_device_name("has space"));
1053 assert!(!is_valid_device_name("../escape"));
1054 assert!(!is_valid_device_name("slash/inside"));
1055 assert!(!is_valid_device_name("querylike?x=1"));
1056 assert!(!is_valid_device_name(&"x".repeat(65)));
1057 }
1058
1059 #[test]
1060 fn constant_time_eq_matches_equality() {
1061 assert!(constant_time_eq("abc", "abc"));
1062 assert!(!constant_time_eq("abc", "abd"));
1063 assert!(!constant_time_eq("abc", "ab"));
1064 assert!(constant_time_eq("", ""));
1065 }
1066}