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 ConnectInfo, FromRequestParts, Request, State,
30 },
31 http::{HeaderMap, StatusCode},
32 response::{IntoResponse, Response},
33 routing::{any, get},
34 Extension, Router,
35};
36use futures_util::{SinkExt, StreamExt};
37
38use crate::error::ShellTunnelError;
39use crate::security::{
40 generate_api_key, rate_limit_middleware, RateLimitCharge, RateLimitConfig, RateLimiter,
41};
42use protocol::{reject, DeviceMessage, RelayMessage, PROTOCOL_VERSION};
43use proxy::{
44 is_forwardable, split_device_path, ProxyRequest, ProxyResponse, POOL_WAIT, REQUEST_TIMEOUT,
45};
46use registry::{Device, DeviceRegistry};
47
48pub use registry::{DeviceRegistry as Registry, POOL_TARGET};
49
50pub const HEARTBEAT_TIMEOUT: Duration = Duration::from_secs(90);
52
53const ENROLL_TIMEOUT: Duration = Duration::from_secs(10);
55
56pub const MAX_RELAY_FRAME: usize = 16 * 1024 * 1024;
70
71#[derive(Debug, Clone)]
73pub struct RelayConfig {
74 pub bind: SocketAddr,
76 pub enroll_token: String,
78 pub rate_limit: RateLimitConfig,
85 #[cfg(feature = "tls")]
87 pub tls: Option<crate::tls::TlsFiles>,
88 pub public_base: Option<String>,
94}
95
96impl RelayConfig {
97 pub fn new(bind: SocketAddr, enroll_token: impl Into<String>) -> Self {
99 Self {
100 bind,
101 enroll_token: enroll_token.into(),
102 rate_limit: RateLimitConfig::default(),
103 #[cfg(feature = "tls")]
104 tls: None,
105 public_base: None,
106 }
107 }
108
109 #[cfg(feature = "tls")]
111 pub fn with_tls(mut self, files: crate::tls::TlsFiles) -> Self {
112 self.tls = Some(files);
113 self
114 }
115
116 pub fn without_rate_limit(mut self) -> Self {
118 self.rate_limit.enabled = false;
119 self
120 }
121
122 pub fn with_public_base(mut self, base: impl Into<String>) -> Self {
124 self.public_base = Some(base.into().trim_end_matches('/').to_string());
125 self
126 }
127
128 pub fn resolved_public_base(&self) -> Option<String> {
139 self.public_base.as_deref().map(|base| {
140 public_base_port_hint(base, self.bind.port()).unwrap_or_else(|| base.to_string())
141 })
142 }
143
144 pub fn public_base_or(&self, observed: Option<String>) -> String {
150 self.resolved_public_base()
151 .or(observed)
152 .unwrap_or_else(|| format!("http://{}", self.bind))
153 }
154
155 pub fn public_url_for(&self, device_id: &str, observed: Option<String>) -> String {
157 format!("{}/d/{}", self.public_base_or(observed), device_id)
158 }
159}
160
161pub fn public_base_port_hint(base: &str, listen_port: u16) -> Option<String> {
171 let (scheme, rest) = base.split_once("://")?;
172 let default_port: u16 = match scheme {
173 "https" => 443,
174 "http" => 80,
175 _ => return None,
176 };
177 if listen_port == default_port {
178 return None;
179 }
180 let authority_end = rest.find('/').unwrap_or(rest.len());
181 let authority = &rest[..authority_end];
182 let has_port = match authority.rfind(']') {
185 Some(bracket) => authority[bracket..].contains(':'),
186 None => authority.contains(':'),
187 };
188 if has_port {
189 return None;
190 }
191 Some(format!(
192 "{scheme}://{authority}:{listen_port}{}",
193 &rest[authority_end..]
194 ))
195}
196
197#[derive(Debug, Clone)]
199pub struct RelayState {
200 config: RelayConfig,
201 devices: DeviceRegistry,
202 limiter: Arc<RateLimiter>,
209}
210
211impl RelayState {
212 pub fn new(config: RelayConfig) -> Self {
214 let limiter = Arc::new(RateLimiter::new(config.rate_limit.clone()));
215 Self {
216 config,
217 devices: DeviceRegistry::new(),
218 limiter,
219 }
220 }
221
222 pub fn devices(&self) -> &DeviceRegistry {
224 &self.devices
225 }
226}
227
228pub fn relay_router(state: RelayState) -> Router {
252 let limiter = Arc::clone(&state.limiter);
253
254 Router::new()
255 .route("/health", get(|| async { "OK" }))
256 .route("/relay/v1/control", get(control_handler))
257 .route("/relay/v1/data", get(data_handler))
258 .route("/relay/v1/devices", get(devices_handler))
259 .route("/d/{*rest}", any(proxy_handler))
260 .layer(axum::middleware::from_fn_with_state(
261 limiter,
262 rate_limit_middleware,
263 ))
264 .with_state(state)
265}
266
267pub async fn bind_relay(config: &RelayConfig) -> crate::Result<tokio::net::TcpListener> {
276 tokio::net::TcpListener::bind(config.bind)
277 .await
278 .map_err(crate::error::ShellTunnelError::Io)
279}
280
281pub async fn serve_relay(config: RelayConfig) -> crate::Result<()> {
287 let listener = bind_relay(&config).await?;
288 serve_relay_on(listener, config).await
289}
290
291pub async fn serve_relay_on(
293 listener: tokio::net::TcpListener,
294 config: RelayConfig,
295) -> crate::Result<()> {
296 let bind = config.bind;
297 #[cfg(feature = "tls")]
298 let tls = config.tls.clone();
299 let state = RelayState::new(config);
300 let router = relay_router(state.clone());
301
302 let sweeper = state.devices().clone();
305 tokio::spawn(async move {
306 let mut ticker = tokio::time::interval(HEARTBEAT_TIMEOUT / 3);
307 loop {
308 ticker.tick().await;
309 for id in sweeper.evict_stale(HEARTBEAT_TIMEOUT) {
310 tracing::info!(target: "relay", device_id = %id, "device evicted (no heartbeat)");
311 }
312 }
313 });
314
315 tracing::info!("relay listening on {}", bind);
318
319 let service = router.into_make_service_with_connect_info::<SocketAddr>();
322
323 #[cfg(feature = "tls")]
324 if let Some(files) = tls {
325 let config = crate::tls::acceptor(files.load()?);
328 crate::tls::watch(files, config.clone());
330 let std_listener = listener.into_std().map_err(ShellTunnelError::Io)?;
331 return axum_server::from_tcp_rustls(std_listener, config)
332 .map_err(ShellTunnelError::Io)?
333 .serve(service)
334 .await
335 .map_err(|e| ShellTunnelError::Io(std::io::Error::other(e.to_string())));
336 }
337
338 axum::serve(listener, service)
339 .await
340 .map_err(|e| ShellTunnelError::Io(std::io::Error::other(e.to_string())))?;
341 Ok(())
342}
343
344fn is_valid_device_name(name: &str) -> bool {
350 !name.is_empty()
351 && name.len() <= 64
352 && name
353 .chars()
354 .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_')
355}
356
357fn observed_base(headers: &HeaderMap, tls: bool) -> Option<String> {
363 let host = headers
364 .get("x-forwarded-host")
365 .or_else(|| headers.get(axum::http::header::HOST))
366 .and_then(|value| value.to_str().ok())?;
367 if host.is_empty() {
368 return None;
369 }
370 let scheme = headers
374 .get("x-forwarded-proto")
375 .and_then(|value| value.to_str().ok())
376 .map(|proto| proto.split(',').next().unwrap_or(proto).trim().to_string())
377 .unwrap_or_else(|| if tls { "https" } else { "http" }.to_string());
378 Some(format!("{scheme}://{host}"))
379}
380
381fn serves_tls(_state: &RelayState) -> bool {
383 #[cfg(feature = "tls")]
384 {
385 _state.config.tls.is_some()
386 }
387 #[cfg(not(feature = "tls"))]
388 {
389 false
390 }
391}
392
393async fn control_handler(
395 ws: WebSocketUpgrade,
396 State(state): State<RelayState>,
397 ConnectInfo(peer): ConnectInfo<SocketAddr>,
398 charge: Option<Extension<RateLimitCharge>>,
399 headers: HeaderMap,
400) -> impl IntoResponse {
401 let observed = observed_base(&headers, serves_tls(&state));
402 let charge = charge.map(|Extension(charge)| charge);
403 ws.on_upgrade(move |socket| control_session(socket, state, observed, peer, charge))
404}
405
406async fn control_session(
408 socket: WebSocket,
409 state: RelayState,
410 observed: Option<String>,
411 peer: SocketAddr,
412 charge: Option<RateLimitCharge>,
413) {
414 let (mut sink, mut stream) = socket.split();
415
416 let first = match tokio::time::timeout(ENROLL_TIMEOUT, stream.next()).await {
419 Ok(Some(Ok(Message::Text(text)))) => text,
420 _ => return,
421 };
422
423 let enroll = match serde_json::from_str::<DeviceMessage>(&first) {
424 Ok(DeviceMessage::Enroll {
425 enroll_token,
426 version,
427 label,
428 device_name,
429 }) => (enroll_token, version, label, device_name),
430 _ => {
431 reject_and_close(
432 &mut sink,
433 reject::BAD_HANDSHAKE,
434 "expected an enroll message",
435 )
436 .await;
437 return;
438 }
439 };
440 let (enroll_token, version, label, device_name) = enroll;
441
442 if version != PROTOCOL_VERSION {
443 reject_and_close(
444 &mut sink,
445 reject::UNSUPPORTED_VERSION,
446 &format!("relay speaks protocol version {PROTOCOL_VERSION}"),
447 )
448 .await;
449 return;
450 }
451
452 if !constant_time_eq(&enroll_token, &state.config.enroll_token) {
453 tracing::debug!(target: "relay", "enrollment rejected: bad token");
457 reject_and_close(&mut sink, reject::BAD_TOKEN, "enrollment refused").await;
458 return;
459 }
460
461 if let Some(charge) = charge {
463 state.limiter.refund(peer.ip(), charge);
464 }
465
466 let device_id = match device_name {
471 Some(name) if !is_valid_device_name(&name) => {
472 reject_and_close(
473 &mut sink,
474 reject::BAD_DEVICE_NAME,
475 "device names may use letters, digits, '-' and '_' (1-64 characters)",
476 )
477 .await;
478 return;
479 }
480 Some(name) => name,
486 None => generate_api_key(),
487 };
488 let public_url = state.config.public_url_for(&device_id, observed);
489 let registry::DeviceHandles {
490 device,
491 mut refill_rx,
492 } = state.devices.attach(&device_id, label.clone());
493 tracing::info!(
494 target: "relay",
495 device_id = %device_id,
496 label = label.as_deref().unwrap_or("-"),
497 "device attached"
498 );
499
500 let enrolled = RelayMessage::Enrolled {
501 device_id: device_id.clone(),
502 public_url,
503 };
504 if send_json(&mut sink, &enrolled).await.is_err() {
505 state.devices.detach(&device_id);
506 return;
507 }
508
509 let fill = RelayMessage::OpenData {
511 count: registry::POOL_TARGET,
512 };
513 if send_json(&mut sink, &fill).await.is_err() {
514 state.devices.detach(&device_id);
515 return;
516 }
517
518 loop {
521 tokio::select! {
522 incoming = stream.next() => {
523 let Some(Ok(message)) = incoming else { break };
524 match message {
525 Message::Text(text) => match serde_json::from_str::<DeviceMessage>(&text) {
526 Ok(DeviceMessage::Heartbeat) => {
527 device.touch();
528 if send_json(&mut sink, &RelayMessage::HeartbeatAck).await.is_err() {
529 break;
530 }
531 }
532 _ => continue,
536 },
537 Message::Close(_) => break,
538 _ => continue,
539 }
540 }
541 refill = refill_rx.recv() => {
542 if refill.is_none() {
543 break;
544 }
545 if send_json(&mut sink, &RelayMessage::OpenData { count: 1 }).await.is_err() {
546 break;
547 }
548 }
549 }
550 }
551
552 state.devices.detach(&device_id);
553 tracing::info!(target: "relay", device_id = %device_id, "device detached");
554}
555
556async fn devices_handler(State(state): State<RelayState>, headers: HeaderMap) -> Response {
562 let presented = headers
563 .get(axum::http::header::AUTHORIZATION)
564 .and_then(|value| value.to_str().ok())
565 .and_then(|value| value.strip_prefix("Bearer "))
566 .unwrap_or("");
567 if !constant_time_eq(presented, &state.config.enroll_token) {
568 return StatusCode::UNAUTHORIZED.into_response();
569 }
570
571 let base = state
572 .config
573 .public_base_or(observed_base(&headers, serves_tls(&state)));
574 let devices: Vec<_> = state
575 .devices
576 .list()
577 .into_iter()
578 .map(|device| {
579 let url = format!("{}/d/{}", base, device.id);
580 let mut entry = serde_json::to_value(&device)
584 .unwrap_or_else(|_| serde_json::json!({ "id": device.id, "label": device.label }));
585 if let Some(object) = entry.as_object_mut() {
586 object.insert("public_url".to_string(), serde_json::Value::String(url));
587 }
588 entry
589 })
590 .collect();
591
592 axum::Json(serde_json::json!({ "devices": devices })).into_response()
593}
594
595async fn data_handler(
602 ws: WebSocketUpgrade,
603 State(state): State<RelayState>,
604 ConnectInfo(peer): ConnectInfo<SocketAddr>,
605 charge: Option<Extension<RateLimitCharge>>,
606) -> Response {
607 let charge = charge.map(|Extension(charge)| charge);
608 ws.max_frame_size(MAX_RELAY_FRAME)
613 .on_upgrade(move |socket| attach_data_connection(socket, state, peer, charge))
614}
615
616async fn attach_data_connection(
618 mut socket: WebSocket,
619 state: RelayState,
620 peer: SocketAddr,
621 charge: Option<RateLimitCharge>,
622) {
623 let first = tokio::time::timeout(ENROLL_TIMEOUT, socket.recv()).await;
624 let Ok(Some(Ok(Message::Text(text)))) = first else {
625 let _ = socket.close().await;
626 return;
627 };
628
629 let Ok(DeviceMessage::Attach {
630 device_id,
631 enroll_token,
632 }) = serde_json::from_str::<DeviceMessage>(&text)
633 else {
634 let _ = socket.close().await;
635 return;
636 };
637
638 if !constant_time_eq(&enroll_token, &state.config.enroll_token) {
639 tracing::debug!(target: "relay", "data connection rejected: bad token");
640 let _ = socket.close().await;
641 return;
642 }
643
644 if let Some(charge) = charge {
649 state.limiter.refund(peer.ip(), charge);
650 }
651
652 let Some(device) = state.devices.get(&device_id) else {
653 let _ = socket.close().await;
654 return;
655 };
656
657 if let Some(mut extra) = device.offer(socket).await {
660 let _ = extra.close().await;
661 }
662}
663
664async fn proxy_handler(State(state): State<RelayState>, request: Request) -> Response {
666 let path_and_query = request
667 .uri()
668 .path_and_query()
669 .map(|p| p.as_str().to_string())
670 .unwrap_or_else(|| request.uri().path().to_string());
671
672 let Some((device_id, tail)) = split_device_path(&path_and_query) else {
673 return StatusCode::NOT_FOUND.into_response();
674 };
675
676 let Some(device) = state.devices.get(device_id) else {
677 return (StatusCode::BAD_GATEWAY, "device is not connected").into_response();
680 };
681
682 let method = request.method().to_string();
683 let headers: Vec<(String, String)> = request
684 .headers()
685 .iter()
686 .filter(|(name, _)| is_forwardable(name.as_str()))
687 .filter_map(|(name, value)| {
688 value
689 .to_str()
690 .ok()
691 .map(|v| (name.as_str().to_string(), v.to_string()))
692 })
693 .collect();
694
695 if is_websocket_upgrade(request.headers()) {
701 let (mut parts, _) = request.into_parts();
702 let upgrade = match WebSocketUpgrade::from_request_parts(&mut parts, &state).await {
703 Ok(upgrade) => upgrade,
704 Err(rejection) => return rejection.into_response(),
705 };
706 let proxied = ProxyRequest {
707 method,
708 path: tail,
709 headers,
710 websocket: true,
711 };
712 return upgrade.on_upgrade(move |client| pipe_websocket(client, device, proxied));
713 }
714
715 let body = match axum::body::to_bytes(request.into_body(), MAX_BODY).await {
716 Ok(body) => body,
717 Err(_) => return StatusCode::PAYLOAD_TOO_LARGE.into_response(),
718 };
719
720 let Some(conn) = device.take(POOL_WAIT).await else {
721 return (
724 StatusCode::SERVICE_UNAVAILABLE,
725 [("retry-after", "1")],
726 "no data connection available",
727 )
728 .into_response();
729 };
730
731 let started = std::time::Instant::now();
737 let outcome = tokio::time::timeout(
738 REQUEST_TIMEOUT,
739 forward(
740 conn,
741 ProxyRequest {
742 method,
743 path: tail,
744 headers,
745 websocket: false,
746 },
747 body,
748 ),
749 )
750 .await;
751 device.record_exchange(started.elapsed());
752
753 match outcome {
754 Ok(Ok(response)) => response,
755 Ok(Err(reason)) => {
756 tracing::debug!(target: "relay", device_id = %device.id, reason, "proxy failed");
757 (StatusCode::BAD_GATEWAY, "device did not answer").into_response()
758 }
759 Err(_) => (StatusCode::GATEWAY_TIMEOUT, "device timed out").into_response(),
760 }
761}
762
763fn is_websocket_upgrade(headers: &HeaderMap) -> bool {
765 let header_contains = |name: axum::http::HeaderName, needle: &str| {
766 headers
767 .get(name)
768 .and_then(|value| value.to_str().ok())
769 .is_some_and(|value| value.to_ascii_lowercase().contains(needle))
770 };
771 header_contains(axum::http::header::UPGRADE, "websocket")
772 && header_contains(axum::http::header::CONNECTION, "upgrade")
773}
774
775async fn pipe_websocket(mut client: WebSocket, device: Arc<Device>, request: ProxyRequest) {
781 let Some(mut conn) = device.take(POOL_WAIT).await else {
782 tracing::debug!(target: "relay", device_id = %device.id, "no data connection for websocket");
783 let _ = client.close().await;
784 return;
785 };
786
787 let Ok(header) = serde_json::to_string(&request) else {
788 let _ = client.close().await;
789 return;
790 };
791 if conn.send(Message::Text(header.into())).await.is_err() {
792 let _ = client.close().await;
793 return;
794 }
795
796 let switched = matches!(
799 conn.recv().await,
800 Some(Ok(Message::Text(ref text)))
801 if serde_json::from_str::<ProxyResponse>(text)
802 .map(|response| response.status == 101)
803 .unwrap_or(false)
804 );
805 if !switched {
806 let _ = client.close().await;
807 let _ = conn.close().await;
808 return;
809 }
810
811 loop {
814 tokio::select! {
815 from_client = client.recv() => {
816 match from_client {
817 Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
818 Some(Ok(message)) => {
819 if conn.send(message).await.is_err() {
820 break;
821 }
822 }
823 }
824 }
825 from_device = conn.recv() => {
826 match from_device {
827 Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
828 Some(Ok(message)) => {
829 if client.send(message).await.is_err() {
830 break;
831 }
832 }
833 }
834 }
835 }
836 }
837
838 let _ = client.close().await;
839 let _ = conn.close().await;
840}
841
842const MAX_BODY: usize = 8 * 1024 * 1024;
844
845async fn forward(
850 mut conn: WebSocket,
851 request: ProxyRequest,
852 body: Bytes,
853) -> Result<Response, &'static str> {
854 let header = serde_json::to_string(&request).map_err(|_| "request-encode")?;
855 conn.send(Message::Text(header.into()))
856 .await
857 .map_err(|_| "request-header-send")?;
858 conn.send(Message::Binary(body))
859 .await
860 .map_err(|_| "request-body-send")?;
861
862 let head: ProxyResponse = loop {
863 match conn.recv().await {
864 Some(Ok(Message::Text(text))) => {
865 break serde_json::from_str(&text).map_err(|_| "response-decode")?
866 }
867 Some(Ok(_)) => continue,
868 _ => return Err("response-header-missing"),
869 }
870 };
871
872 let mut body = Vec::new();
881 loop {
882 match conn.recv().await {
883 Some(Ok(Message::Binary(chunk))) => body.extend_from_slice(&chunk),
884 Some(Ok(Message::Close(_))) | None => break,
885 Some(Ok(_)) => continue,
886 Some(Err(_)) => return Err("response-body-truncated"),
887 }
888 }
889
890 let mut response = Response::builder().status(head.status);
891 for (name, value) in head.headers {
892 if is_forwardable(&name) {
893 response = response.header(name, value);
894 }
895 }
896 response
897 .body(axum::body::Body::from(body))
898 .map_err(|_| "response-build")
899}
900
901async fn reject_and_close<S>(sink: &mut S, code: &str, message: &str)
903where
904 S: SinkExt<Message> + Unpin,
905{
906 let rejected = RelayMessage::Rejected {
907 code: code.to_string(),
908 message: message.to_string(),
909 };
910 let _ = send_json(sink, &rejected).await;
911 let _ = sink.close().await;
912}
913
914async fn send_json<S, T>(sink: &mut S, message: &T) -> Result<(), ()>
916where
917 S: SinkExt<Message> + Unpin,
918 T: serde::Serialize,
919{
920 let json = serde_json::to_string(message).map_err(|_| ())?;
921 sink.send(Message::Text(json.into())).await.map_err(|_| ())
922}
923
924fn constant_time_eq(a: &str, b: &str) -> bool {
930 let (a, b) = (a.as_bytes(), b.as_bytes());
931 if a.len() != b.len() {
932 return false;
933 }
934 a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
935}
936
937#[cfg(test)]
938mod tests {
939 use super::*;
940
941 fn config() -> RelayConfig {
942 RelayConfig::new("127.0.0.1:0".parse().unwrap(), "secret")
943 }
944
945 #[test]
946 fn a_portless_base_on_a_nondefault_port_gets_a_corrected_suggestion() {
947 assert_eq!(
953 public_base_port_hint("https://labs.example.com", 8443).as_deref(),
954 Some("https://labs.example.com:8443")
955 );
956 assert_eq!(
957 public_base_port_hint("http://relay.local", 8080).as_deref(),
958 Some("http://relay.local:8080")
959 );
960 }
961
962 #[test]
963 fn a_base_matching_the_scheme_default_needs_no_hint() {
964 assert_eq!(public_base_port_hint("https://labs.example.com", 443), None);
965 assert_eq!(public_base_port_hint("http://relay.local", 80), None);
966 }
967
968 #[test]
969 fn an_explicit_port_is_the_operator_stating_intent() {
970 assert_eq!(
972 public_base_port_hint("https://labs.example.com:8443", 8443),
973 None
974 );
975 assert_eq!(
976 public_base_port_hint("https://labs.example.com:9000", 8443),
977 None
978 );
979 assert_eq!(
980 public_base_port_hint("https://labs.example.com:443", 8443),
981 None
982 );
983 }
984
985 #[test]
986 fn the_port_is_spliced_into_the_authority_not_the_tail() {
987 assert_eq!(
989 public_base_port_hint("https://labs.example.com/relay", 8443).as_deref(),
990 Some("https://labs.example.com:8443/relay")
991 );
992 }
993
994 #[test]
995 fn ipv6_literals_look_for_the_port_after_the_bracket() {
996 assert_eq!(
997 public_base_port_hint("https://[::1]", 8443).as_deref(),
998 Some("https://[::1]:8443")
999 );
1000 assert_eq!(public_base_port_hint("https://[::1]:8443", 8443), None);
1001 }
1002
1003 #[test]
1004 fn an_unrecognized_scheme_is_left_alone() {
1005 assert_eq!(public_base_port_hint("ws://relay.local", 8443), None);
1006 }
1007
1008 #[test]
1009 fn public_url_uses_the_device_path_prefix() {
1010 let config = RelayConfig::new("127.0.0.1:443".parse().unwrap(), "secret")
1013 .with_public_base("https://relay.example.com/");
1014 assert_eq!(
1015 config.public_url_for("dev-1", None),
1016 "https://relay.example.com/d/dev-1"
1017 );
1018 }
1019
1020 #[test]
1021 fn a_portless_base_inherits_the_listen_port() {
1022 let config = RelayConfig::new("0.0.0.0:8443".parse().unwrap(), "secret")
1026 .with_public_base("https://labs.example.com");
1027 assert_eq!(
1028 config.resolved_public_base().as_deref(),
1029 Some("https://labs.example.com:8443")
1030 );
1031 assert_eq!(
1032 config.public_url_for("dev-1", None),
1033 "https://labs.example.com:8443/d/dev-1"
1034 );
1035 }
1036
1037 #[test]
1038 fn an_explicit_port_survives_resolution() {
1039 let config = RelayConfig::new("0.0.0.0:8443".parse().unwrap(), "secret")
1042 .with_public_base("https://labs.example.com:443");
1043 assert_eq!(
1044 config.resolved_public_base().as_deref(),
1045 Some("https://labs.example.com:443")
1046 );
1047 }
1048
1049 #[test]
1050 fn resolution_leaves_a_default_port_base_alone() {
1051 let config = RelayConfig::new("0.0.0.0:443".parse().unwrap(), "secret")
1053 .with_public_base("https://labs.example.com");
1054 assert_eq!(
1055 config.resolved_public_base().as_deref(),
1056 Some("https://labs.example.com")
1057 );
1058 }
1059
1060 #[test]
1061 fn public_base_defaults_to_the_bind_address() {
1062 let config = RelayConfig::new("127.0.0.1:8443".parse().unwrap(), "secret");
1063 assert_eq!(
1064 config.public_url_for("d", None),
1065 "http://127.0.0.1:8443/d/d"
1066 );
1067 }
1068
1069 #[test]
1070 fn an_observed_address_is_used_when_the_operator_configured_none() {
1071 let config = config();
1072 assert_eq!(
1073 config.public_url_for("dev-1", Some("https://relay.example.com".into())),
1074 "https://relay.example.com/d/dev-1"
1075 );
1076 }
1077
1078 #[test]
1079 fn a_configured_base_wins_over_what_the_connection_observed() {
1080 let config = RelayConfig::new("127.0.0.1:443".parse().unwrap(), "secret")
1081 .with_public_base("https://canonical.example");
1082 assert_eq!(
1083 config.public_url_for("dev-1", Some("https://whatever.invalid".into())),
1084 "https://canonical.example/d/dev-1"
1085 );
1086 }
1087
1088 #[test]
1089 fn the_forwarded_scheme_and_host_are_preferred_over_the_direct_host() {
1090 let mut headers = HeaderMap::new();
1091 headers.insert(axum::http::header::HOST, "127.0.0.1:8443".parse().unwrap());
1092 assert_eq!(
1093 observed_base(&headers, false).as_deref(),
1094 Some("http://127.0.0.1:8443")
1095 );
1096
1097 headers.insert("x-forwarded-proto", "https".parse().unwrap());
1098 headers.insert("x-forwarded-host", "relay.example.com".parse().unwrap());
1099 assert_eq!(
1100 observed_base(&headers, false).as_deref(),
1101 Some("https://relay.example.com")
1102 );
1103 }
1104
1105 #[test]
1106 fn a_proxy_chain_scheme_takes_the_first_entry() {
1107 let mut headers = HeaderMap::new();
1108 headers.insert(
1109 axum::http::header::HOST,
1110 "relay.example.com".parse().unwrap(),
1111 );
1112 headers.insert("x-forwarded-proto", "https, http".parse().unwrap());
1113 assert_eq!(
1114 observed_base(&headers, false).as_deref(),
1115 Some("https://relay.example.com")
1116 );
1117 }
1118
1119 #[test]
1120 fn no_host_header_means_nothing_observed() {
1121 assert!(observed_base(&HeaderMap::new(), false).is_none());
1122 }
1123
1124 #[test]
1125 fn terminating_tls_makes_the_advertised_url_https() {
1126 let mut headers = HeaderMap::new();
1129 headers.insert(
1130 axum::http::header::HOST,
1131 "relay.example.com".parse().unwrap(),
1132 );
1133 assert_eq!(
1134 observed_base(&headers, true).as_deref(),
1135 Some("https://relay.example.com")
1136 );
1137 }
1138
1139 #[test]
1140 fn device_names_must_be_url_path_safe() {
1141 assert!(is_valid_device_name("build-box"));
1142 assert!(is_valid_device_name("laptop_2"));
1143 assert!(is_valid_device_name("a"));
1144
1145 assert!(!is_valid_device_name(""));
1146 assert!(!is_valid_device_name("has space"));
1147 assert!(!is_valid_device_name("../escape"));
1148 assert!(!is_valid_device_name("slash/inside"));
1149 assert!(!is_valid_device_name("querylike?x=1"));
1150 assert!(!is_valid_device_name(&"x".repeat(65)));
1151 }
1152
1153 #[test]
1154 fn constant_time_eq_matches_equality() {
1155 assert!(constant_time_eq("abc", "abc"));
1156 assert!(!constant_time_eq("abc", "abd"));
1157 assert!(!constant_time_eq("abc", "ab"));
1158 assert!(constant_time_eq("", ""));
1159 }
1160}