#[cfg(feature = "relay-client")]
pub mod client;
pub mod protocol;
pub mod proxy;
pub mod registry;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use axum::{
body::Bytes,
extract::{
ws::{Message, WebSocket, WebSocketUpgrade},
FromRequestParts, Request, State,
},
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
routing::{any, get},
Router,
};
use futures_util::{SinkExt, StreamExt};
use crate::error::ShellTunnelError;
use crate::security::generate_api_key;
use protocol::{reject, DeviceMessage, RelayMessage, PROTOCOL_VERSION};
use proxy::{
is_forwardable, split_device_path, ProxyRequest, ProxyResponse, POOL_WAIT, REQUEST_TIMEOUT,
};
use registry::{Device, DeviceRegistry};
pub use registry::{DeviceRegistry as Registry, POOL_TARGET};
pub const HEARTBEAT_TIMEOUT: Duration = Duration::from_secs(90);
const ENROLL_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Debug, Clone)]
pub struct RelayConfig {
pub bind: SocketAddr,
pub enroll_token: String,
pub public_base: Option<String>,
}
impl RelayConfig {
pub fn new(bind: SocketAddr, enroll_token: impl Into<String>) -> Self {
Self {
bind,
enroll_token: enroll_token.into(),
public_base: None,
}
}
pub fn with_public_base(mut self, base: impl Into<String>) -> Self {
self.public_base = Some(base.into().trim_end_matches('/').to_string());
self
}
pub fn public_base_or(&self, observed: Option<String>) -> String {
self.public_base
.clone()
.or(observed)
.unwrap_or_else(|| format!("http://{}", self.bind))
}
pub fn public_url_for(&self, device_id: &str, observed: Option<String>) -> String {
format!("{}/d/{}", self.public_base_or(observed), device_id)
}
}
#[derive(Debug, Clone)]
pub struct RelayState {
config: RelayConfig,
devices: DeviceRegistry,
}
impl RelayState {
pub fn new(config: RelayConfig) -> Self {
Self {
config,
devices: DeviceRegistry::new(),
}
}
pub fn devices(&self) -> &DeviceRegistry {
&self.devices
}
}
pub fn relay_router(state: RelayState) -> Router {
Router::new()
.route("/health", get(|| async { "OK" }))
.route("/relay/v1/control", get(control_handler))
.route("/relay/v1/data", get(data_handler))
.route("/d/{*rest}", any(proxy_handler))
.with_state(state)
}
pub async fn serve_relay(config: RelayConfig) -> crate::Result<()> {
let bind = config.bind;
let state = RelayState::new(config);
let router = relay_router(state.clone());
let sweeper = state.devices().clone();
tokio::spawn(async move {
let mut ticker = tokio::time::interval(HEARTBEAT_TIMEOUT / 3);
loop {
ticker.tick().await;
for id in sweeper.evict_stale(HEARTBEAT_TIMEOUT) {
tracing::info!(target: "relay", device_id = %id, "device evicted (no heartbeat)");
}
}
});
tracing::info!("relay listening on {}", bind);
let listener = tokio::net::TcpListener::bind(bind)
.await
.map_err(ShellTunnelError::Io)?;
axum::serve(listener, router)
.await
.map_err(|e| ShellTunnelError::Io(std::io::Error::other(e.to_string())))?;
Ok(())
}
fn observed_base(headers: &HeaderMap) -> Option<String> {
let host = headers
.get("x-forwarded-host")
.or_else(|| headers.get(axum::http::header::HOST))
.and_then(|value| value.to_str().ok())?;
if host.is_empty() {
return None;
}
let scheme = headers
.get("x-forwarded-proto")
.and_then(|value| value.to_str().ok())
.map(|proto| proto.split(',').next().unwrap_or(proto).trim().to_string())
.unwrap_or_else(|| "http".to_string());
Some(format!("{scheme}://{host}"))
}
async fn control_handler(
ws: WebSocketUpgrade,
State(state): State<RelayState>,
headers: HeaderMap,
) -> impl IntoResponse {
let observed = observed_base(&headers);
ws.on_upgrade(move |socket| control_session(socket, state, observed))
}
async fn control_session(socket: WebSocket, state: RelayState, observed: Option<String>) {
let (mut sink, mut stream) = socket.split();
let first = match tokio::time::timeout(ENROLL_TIMEOUT, stream.next()).await {
Ok(Some(Ok(Message::Text(text)))) => text,
_ => return,
};
let enroll = match serde_json::from_str::<DeviceMessage>(&first) {
Ok(DeviceMessage::Enroll {
enroll_token,
version,
label,
}) => (enroll_token, version, label),
_ => {
reject_and_close(
&mut sink,
reject::BAD_HANDSHAKE,
"expected an enroll message",
)
.await;
return;
}
};
let (enroll_token, version, label) = enroll;
if version != PROTOCOL_VERSION {
reject_and_close(
&mut sink,
reject::UNSUPPORTED_VERSION,
&format!("relay speaks protocol version {PROTOCOL_VERSION}"),
)
.await;
return;
}
if !constant_time_eq(&enroll_token, &state.config.enroll_token) {
tracing::debug!(target: "relay", "enrollment rejected: bad token");
reject_and_close(&mut sink, reject::BAD_TOKEN, "enrollment refused").await;
return;
}
let device_id = generate_api_key();
let public_url = state.config.public_url_for(&device_id, observed);
let registry::DeviceHandles {
device,
mut refill_rx,
} = state.devices.attach(&device_id, label.clone());
tracing::info!(
target: "relay",
device_id = %device_id,
label = label.as_deref().unwrap_or("-"),
"device attached"
);
let enrolled = RelayMessage::Enrolled {
device_id: device_id.clone(),
public_url,
};
if send_json(&mut sink, &enrolled).await.is_err() {
state.devices.detach(&device_id);
return;
}
let fill = RelayMessage::OpenData {
count: registry::POOL_TARGET,
};
if send_json(&mut sink, &fill).await.is_err() {
state.devices.detach(&device_id);
return;
}
loop {
tokio::select! {
incoming = stream.next() => {
let Some(Ok(message)) = incoming else { break };
match message {
Message::Text(text) => match serde_json::from_str::<DeviceMessage>(&text) {
Ok(DeviceMessage::Heartbeat) => {
device.touch();
if send_json(&mut sink, &RelayMessage::HeartbeatAck).await.is_err() {
break;
}
}
_ => continue,
},
Message::Close(_) => break,
_ => continue,
}
}
refill = refill_rx.recv() => {
if refill.is_none() {
break;
}
if send_json(&mut sink, &RelayMessage::OpenData { count: 1 }).await.is_err() {
break;
}
}
}
}
state.devices.detach(&device_id);
tracing::info!(target: "relay", device_id = %device_id, "device detached");
}
async fn data_handler(ws: WebSocketUpgrade, State(state): State<RelayState>) -> Response {
ws.on_upgrade(move |socket| attach_data_connection(socket, state))
}
async fn attach_data_connection(mut socket: WebSocket, state: RelayState) {
let first = tokio::time::timeout(ENROLL_TIMEOUT, socket.recv()).await;
let Ok(Some(Ok(Message::Text(text)))) = first else {
let _ = socket.close().await;
return;
};
let Ok(DeviceMessage::Attach {
device_id,
enroll_token,
}) = serde_json::from_str::<DeviceMessage>(&text)
else {
let _ = socket.close().await;
return;
};
if !constant_time_eq(&enroll_token, &state.config.enroll_token) {
tracing::debug!(target: "relay", "data connection rejected: bad token");
let _ = socket.close().await;
return;
}
let Some(device) = state.devices.get(&device_id) else {
let _ = socket.close().await;
return;
};
if let Some(mut extra) = device.offer(socket).await {
let _ = extra.close().await;
}
}
async fn proxy_handler(State(state): State<RelayState>, request: Request) -> Response {
let path_and_query = request
.uri()
.path_and_query()
.map(|p| p.as_str().to_string())
.unwrap_or_else(|| request.uri().path().to_string());
let Some((device_id, tail)) = split_device_path(&path_and_query) else {
return StatusCode::NOT_FOUND.into_response();
};
let Some(device) = state.devices.get(device_id) else {
return (StatusCode::BAD_GATEWAY, "device is not connected").into_response();
};
let method = request.method().to_string();
let headers: Vec<(String, String)> = request
.headers()
.iter()
.filter(|(name, _)| is_forwardable(name.as_str()))
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|v| (name.as_str().to_string(), v.to_string()))
})
.collect();
if is_websocket_upgrade(request.headers()) {
let (mut parts, _) = request.into_parts();
let upgrade = match WebSocketUpgrade::from_request_parts(&mut parts, &state).await {
Ok(upgrade) => upgrade,
Err(rejection) => return rejection.into_response(),
};
let proxied = ProxyRequest {
method,
path: tail,
headers,
websocket: true,
};
return upgrade.on_upgrade(move |client| pipe_websocket(client, device, proxied));
}
let body = match axum::body::to_bytes(request.into_body(), MAX_BODY).await {
Ok(body) => body,
Err(_) => return StatusCode::PAYLOAD_TOO_LARGE.into_response(),
};
let Some(conn) = device.take(POOL_WAIT).await else {
return (
StatusCode::SERVICE_UNAVAILABLE,
[("retry-after", "1")],
"no data connection available",
)
.into_response();
};
match tokio::time::timeout(
REQUEST_TIMEOUT,
forward(
conn,
ProxyRequest {
method,
path: tail,
headers,
websocket: false,
},
body,
),
)
.await
{
Ok(Ok(response)) => response,
Ok(Err(reason)) => {
tracing::debug!(target: "relay", device_id = %device.id, reason, "proxy failed");
(StatusCode::BAD_GATEWAY, "device did not answer").into_response()
}
Err(_) => (StatusCode::GATEWAY_TIMEOUT, "device timed out").into_response(),
}
}
fn is_websocket_upgrade(headers: &HeaderMap) -> bool {
let header_contains = |name: axum::http::HeaderName, needle: &str| {
headers
.get(name)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.to_ascii_lowercase().contains(needle))
};
header_contains(axum::http::header::UPGRADE, "websocket")
&& header_contains(axum::http::header::CONNECTION, "upgrade")
}
async fn pipe_websocket(mut client: WebSocket, device: Arc<Device>, request: ProxyRequest) {
let Some(mut conn) = device.take(POOL_WAIT).await else {
tracing::debug!(target: "relay", device_id = %device.id, "no data connection for websocket");
let _ = client.close().await;
return;
};
let Ok(header) = serde_json::to_string(&request) else {
let _ = client.close().await;
return;
};
if conn.send(Message::Text(header.into())).await.is_err() {
let _ = client.close().await;
return;
}
let switched = matches!(
conn.recv().await,
Some(Ok(Message::Text(ref text)))
if serde_json::from_str::<ProxyResponse>(text)
.map(|response| response.status == 101)
.unwrap_or(false)
);
if !switched {
let _ = client.close().await;
let _ = conn.close().await;
return;
}
loop {
tokio::select! {
from_client = client.recv() => {
match from_client {
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
Some(Ok(message)) => {
if conn.send(message).await.is_err() {
break;
}
}
}
}
from_device = conn.recv() => {
match from_device {
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
Some(Ok(message)) => {
if client.send(message).await.is_err() {
break;
}
}
}
}
}
}
let _ = client.close().await;
let _ = conn.close().await;
}
const MAX_BODY: usize = 8 * 1024 * 1024;
async fn forward(
mut conn: WebSocket,
request: ProxyRequest,
body: Bytes,
) -> Result<Response, &'static str> {
let header = serde_json::to_string(&request).map_err(|_| "request-encode")?;
conn.send(Message::Text(header.into()))
.await
.map_err(|_| "request-header-send")?;
conn.send(Message::Binary(body))
.await
.map_err(|_| "request-body-send")?;
let head: ProxyResponse = loop {
match conn.recv().await {
Some(Ok(Message::Text(text))) => {
break serde_json::from_str(&text).map_err(|_| "response-decode")?
}
Some(Ok(_)) => continue,
_ => return Err("response-header-missing"),
}
};
let mut body = Vec::new();
while let Some(Ok(message)) = conn.recv().await {
match message {
Message::Binary(chunk) => body.extend_from_slice(&chunk),
Message::Close(_) => break,
_ => continue,
}
}
let mut response = Response::builder().status(head.status);
for (name, value) in head.headers {
if is_forwardable(&name) {
response = response.header(name, value);
}
}
response
.body(axum::body::Body::from(body))
.map_err(|_| "response-build")
}
async fn reject_and_close<S>(sink: &mut S, code: &str, message: &str)
where
S: SinkExt<Message> + Unpin,
{
let rejected = RelayMessage::Rejected {
code: code.to_string(),
message: message.to_string(),
};
let _ = send_json(sink, &rejected).await;
let _ = sink.close().await;
}
async fn send_json<S, T>(sink: &mut S, message: &T) -> Result<(), ()>
where
S: SinkExt<Message> + Unpin,
T: serde::Serialize,
{
let json = serde_json::to_string(message).map_err(|_| ())?;
sink.send(Message::Text(json.into())).await.map_err(|_| ())
}
fn constant_time_eq(a: &str, b: &str) -> bool {
let (a, b) = (a.as_bytes(), b.as_bytes());
if a.len() != b.len() {
return false;
}
a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
}
#[cfg(test)]
mod tests {
use super::*;
fn config() -> RelayConfig {
RelayConfig::new("127.0.0.1:0".parse().unwrap(), "secret")
}
#[test]
fn public_url_uses_the_device_path_prefix() {
let config = config().with_public_base("https://relay.example.com/");
assert_eq!(
config.public_url_for("dev-1", None),
"https://relay.example.com/d/dev-1"
);
}
#[test]
fn public_base_defaults_to_the_bind_address() {
let config = RelayConfig::new("127.0.0.1:8443".parse().unwrap(), "secret");
assert_eq!(
config.public_url_for("d", None),
"http://127.0.0.1:8443/d/d"
);
}
#[test]
fn an_observed_address_is_used_when_the_operator_configured_none() {
let config = config();
assert_eq!(
config.public_url_for("dev-1", Some("https://relay.example.com".into())),
"https://relay.example.com/d/dev-1"
);
}
#[test]
fn a_configured_base_wins_over_what_the_connection_observed() {
let config = config().with_public_base("https://canonical.example");
assert_eq!(
config.public_url_for("dev-1", Some("https://whatever.invalid".into())),
"https://canonical.example/d/dev-1"
);
}
#[test]
fn the_forwarded_scheme_and_host_are_preferred_over_the_direct_host() {
let mut headers = HeaderMap::new();
headers.insert(axum::http::header::HOST, "127.0.0.1:8443".parse().unwrap());
assert_eq!(
observed_base(&headers).as_deref(),
Some("http://127.0.0.1:8443")
);
headers.insert("x-forwarded-proto", "https".parse().unwrap());
headers.insert("x-forwarded-host", "relay.example.com".parse().unwrap());
assert_eq!(
observed_base(&headers).as_deref(),
Some("https://relay.example.com")
);
}
#[test]
fn a_proxy_chain_scheme_takes_the_first_entry() {
let mut headers = HeaderMap::new();
headers.insert(
axum::http::header::HOST,
"relay.example.com".parse().unwrap(),
);
headers.insert("x-forwarded-proto", "https, http".parse().unwrap());
assert_eq!(
observed_base(&headers).as_deref(),
Some("https://relay.example.com")
);
}
#[test]
fn no_host_header_means_nothing_observed() {
assert!(observed_base(&HeaderMap::new()).is_none());
}
#[test]
fn constant_time_eq_matches_equality() {
assert!(constant_time_eq("abc", "abc"));
assert!(!constant_time_eq("abc", "abd"));
assert!(!constant_time_eq("abc", "ab"));
assert!(constant_time_eq("", ""));
}
}