use std::net::SocketAddr;
use std::time::Duration;
use futures_util::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::{HeaderName, HeaderValue};
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
type WsStream = WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>;
use super::protocol::{DeviceMessage, RelayMessage, PROTOCOL_VERSION};
use super::proxy::{is_forwardable, ProxyRequest, ProxyResponse};
use crate::error::ShellTunnelError;
use crate::Result;
const HEARTBEAT_INTERVAL: Duration = Duration::from_secs(30);
const BACKOFF_MIN: Duration = Duration::from_secs(1);
const BACKOFF_MAX: Duration = Duration::from_secs(60);
#[derive(Debug, Clone)]
pub struct RelayClientConfig {
pub relay_url: String,
pub enroll_token: String,
pub local: SocketAddr,
pub label: Option<String>,
}
impl RelayClientConfig {
pub fn control_url(&self) -> String {
format!("{}/relay/v1/control", self.base())
}
pub fn data_url(&self) -> String {
format!("{}/relay/v1/data", self.base())
}
fn base(&self) -> String {
let trimmed = self.relay_url.trim_end_matches('/');
match trimmed.split_once("://") {
Some(("https", rest)) => format!("wss://{rest}"),
Some(("http", rest)) => format!("ws://{rest}"),
Some(_) => trimmed.to_string(),
None => format!("wss://{trimmed}"),
}
}
}
fn install_crypto_provider() {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
let _ = rustls::crypto::ring::default_provider().install_default();
});
}
pub async fn run(config: RelayClientConfig) -> Result<()> {
install_crypto_provider();
let mut backoff = BACKOFF_MIN;
loop {
match attach(&config).await {
Ok(()) => {
tracing::warn!(target: "relay-client", "relay connection closed; reconnecting");
backoff = BACKOFF_MIN;
}
Err(e) => {
tracing::warn!(target: "relay-client", "relay connection failed: {e}");
}
}
tokio::time::sleep(backoff).await;
backoff = (backoff * 2).min(BACKOFF_MAX);
}
}
pub async fn attach(config: &RelayClientConfig) -> Result<()> {
install_crypto_provider();
let (mut control, _) = tokio_tungstenite::connect_async(config.control_url())
.await
.map_err(|e| ShellTunnelError::Tunnel(format!("cannot reach relay: {e}")))?;
let enroll = DeviceMessage::Enroll {
enroll_token: config.enroll_token.clone(),
version: PROTOCOL_VERSION,
label: config.label.clone(),
};
send(&mut control, &enroll).await?;
let device_id = match recv(&mut control).await? {
RelayMessage::Enrolled {
device_id,
public_url,
} => {
println!("\nPublic URL: {public_url} (via relay)");
device_id
}
RelayMessage::Rejected { code, message } => {
return Err(ShellTunnelError::Tunnel(format!(
"relay refused this device ({code}): {message}"
)))
}
other => {
return Err(ShellTunnelError::Tunnel(format!(
"unexpected first message from relay: {other:?}"
)))
}
};
let mut heartbeat = tokio::time::interval(HEARTBEAT_INTERVAL);
heartbeat.tick().await;
loop {
tokio::select! {
incoming = control.next() => {
let Some(Ok(message)) = incoming else { return Ok(()) };
let Message::Text(text) = message else { continue };
match serde_json::from_str::<RelayMessage>(&text) {
Ok(RelayMessage::OpenData { count }) => {
for _ in 0..count {
spawn_data_connection(config.clone(), device_id.clone());
}
}
Ok(RelayMessage::HeartbeatAck) => {}
_ => continue,
}
}
_ = heartbeat.tick() => {
send(&mut control, &DeviceMessage::Heartbeat).await?;
}
}
}
}
fn spawn_data_connection(config: RelayClientConfig, device_id: String) {
tokio::spawn(async move {
if let Err(e) = serve_one(&config, &device_id).await {
tracing::debug!(target: "relay-client", "data connection ended: {e}");
}
});
}
async fn serve_one(config: &RelayClientConfig, device_id: &str) -> Result<()> {
let (mut conn, _) = tokio_tungstenite::connect_async(config.data_url())
.await
.map_err(|e| ShellTunnelError::Tunnel(format!("data connection refused: {e}")))?;
let attach = DeviceMessage::Attach {
device_id: device_id.to_string(),
enroll_token: config.enroll_token.clone(),
};
send(&mut conn, &attach).await?;
let request: ProxyRequest = loop {
match conn.next().await {
Some(Ok(Message::Text(text))) => {
break serde_json::from_str(&text)
.map_err(|e| ShellTunnelError::Tunnel(format!("bad request header: {e}")))?
}
Some(Ok(_)) => continue,
_ => return Ok(()), }
};
if request.websocket {
return pipe_websocket(conn, config, &request).await;
}
let body = match conn.next().await {
Some(Ok(Message::Binary(bytes))) => bytes.to_vec(),
_ => Vec::new(),
};
let (status, headers, body) = replay_locally(config.local, &request, body).await;
let head = ProxyResponse { status, headers };
let json = serde_json::to_string(&head)
.map_err(|e| ShellTunnelError::Tunnel(format!("cannot encode response: {e}")))?;
let _ = conn.send(Message::Text(json)).await;
let _ = conn.send(Message::Binary(body)).await;
let _ = conn.close(None).await;
Ok(())
}
async fn pipe_websocket(
mut conn: WsStream,
config: &RelayClientConfig,
request: &ProxyRequest,
) -> Result<()> {
let local_url = format!("ws://{}{}", config.local, request.path);
let mut builder = local_url
.into_client_request()
.map_err(|e| ShellTunnelError::Tunnel(format!("bad local websocket url: {e}")))?;
for (name, value) in &request.headers {
if !is_forwardable(name) || name.eq_ignore_ascii_case("sec-websocket-key") {
continue;
}
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(name.as_bytes()),
HeaderValue::from_str(value),
) {
builder.headers_mut().insert(name, value);
}
}
let local = match tokio_tungstenite::connect_async(builder).await {
Ok((socket, _)) => socket,
Err(e) => {
tracing::debug!(target: "relay-client", "local websocket refused: {e}");
let head = ProxyResponse {
status: 502,
headers: Vec::new(),
};
if let Ok(json) = serde_json::to_string(&head) {
let _ = conn.send(Message::Text(json)).await;
}
let _ = conn.close(None).await;
return Ok(());
}
};
let head = ProxyResponse {
status: 101,
headers: Vec::new(),
};
let json = serde_json::to_string(&head)
.map_err(|e| ShellTunnelError::Tunnel(format!("cannot encode response: {e}")))?;
conn.send(Message::Text(json))
.await
.map_err(|_| ShellTunnelError::Tunnel("relay connection lost".to_string()))?;
let (mut local_tx, mut local_rx) = local.split();
let (mut relay_tx, mut relay_rx) = conn.split();
loop {
tokio::select! {
from_relay = relay_rx.next() => {
match from_relay {
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
Some(Ok(message)) => {
if local_tx.send(message).await.is_err() {
break;
}
}
}
}
from_local = local_rx.next() => {
match from_local {
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
Some(Ok(message)) => {
if relay_tx.send(message).await.is_err() {
break;
}
}
}
}
}
}
let _ = local_tx.close().await;
let _ = relay_tx.close().await;
Ok(())
}
async fn replay_locally(
local: SocketAddr,
request: &ProxyRequest,
body: Vec<u8>,
) -> (u16, Vec<(String, String)>, Vec<u8>) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut stream = match tokio::net::TcpStream::connect(local).await {
Ok(stream) => stream,
Err(e) => return bad_gateway(format!("local server unreachable: {e}")),
};
let mut head = format!(
"{} {} HTTP/1.1\r\nHost: {}\r\nConnection: close\r\ncontent-length: {}\r\n",
request.method,
request.path,
local,
body.len()
);
for (name, value) in &request.headers {
if is_forwardable(name) && !name.eq_ignore_ascii_case("content-length") {
head.push_str(&format!("{name}: {value}\r\n"));
}
}
head.push_str("\r\n");
if stream.write_all(head.as_bytes()).await.is_err() || stream.write_all(&body).await.is_err() {
return bad_gateway("local server closed the connection".to_string());
}
let mut raw = Vec::new();
if stream.read_to_end(&mut raw).await.is_err() {
return bad_gateway("local server response was cut short".to_string());
}
parse_response(&raw)
}
fn parse_response(raw: &[u8]) -> (u16, Vec<(String, String)>, Vec<u8>) {
let split = raw
.windows(4)
.position(|w| w == b"\r\n\r\n")
.map(|i| i + 4)
.unwrap_or(raw.len());
let (head, body) = raw.split_at(split);
let head = String::from_utf8_lossy(head);
let mut lines = head.lines();
let status = lines
.next()
.and_then(|line| line.split_whitespace().nth(1))
.and_then(|code| code.parse().ok())
.unwrap_or(502);
let headers = lines
.filter_map(|line| line.split_once(':'))
.map(|(name, value)| (name.trim().to_string(), value.trim().to_string()))
.filter(|(name, _)| is_forwardable(name))
.collect();
(status, headers, body.to_vec())
}
fn bad_gateway(reason: String) -> (u16, Vec<(String, String)>, Vec<u8>) {
tracing::debug!(target: "relay-client", "{reason}");
(
502,
vec![("content-type".to_string(), "text/plain".to_string())],
b"device could not reach its local server".to_vec(),
)
}
async fn send<S>(socket: &mut S, message: &DeviceMessage) -> Result<()>
where
S: SinkExt<Message> + Unpin,
{
let json = serde_json::to_string(message)
.map_err(|e| ShellTunnelError::Tunnel(format!("cannot encode message: {e}")))?;
socket
.send(Message::Text(json))
.await
.map_err(|_| ShellTunnelError::Tunnel("relay connection lost".to_string()))
}
async fn recv<S>(socket: &mut S) -> Result<RelayMessage>
where
S: StreamExt<Item = std::result::Result<Message, tokio_tungstenite::tungstenite::Error>>
+ Unpin,
{
loop {
match socket.next().await {
Some(Ok(Message::Text(text))) => {
return serde_json::from_str(&text)
.map_err(|e| ShellTunnelError::Tunnel(format!("bad relay message: {e}")))
}
Some(Ok(_)) => continue,
_ => {
return Err(ShellTunnelError::Tunnel(
"relay closed the connection".to_string(),
))
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn config(relay_url: &str) -> RelayClientConfig {
RelayClientConfig {
relay_url: relay_url.to_string(),
enroll_token: "secret".to_string(),
local: "127.0.0.1:3000".parse().unwrap(),
label: None,
}
}
#[test]
fn https_urls_become_websocket_urls() {
assert_eq!(
config("https://relay.example.com").control_url(),
"wss://relay.example.com/relay/v1/control"
);
assert_eq!(
config("http://127.0.0.1:8443").control_url(),
"ws://127.0.0.1:8443/relay/v1/control"
);
}
#[test]
fn websocket_urls_are_left_alone() {
assert_eq!(
config("wss://relay.example.com/").control_url(),
"wss://relay.example.com/relay/v1/control"
);
}
#[test]
fn a_bare_host_defaults_to_the_secure_scheme() {
assert_eq!(
config("relay.example.com").control_url(),
"wss://relay.example.com/relay/v1/control"
);
}
#[test]
fn data_urls_carry_no_credentials() {
let url = config("wss://relay.example.com").data_url();
assert_eq!(url, "wss://relay.example.com/relay/v1/data");
assert!(!url.contains("secret"), "{url}");
assert!(!url.contains('?'), "{url}");
}
#[test]
fn responses_are_split_into_status_headers_and_body() {
let raw = b"HTTP/1.1 201 Created\r\ncontent-type: application/json\r\nconnection: close\r\n\r\n{\"ok\":true}";
let (status, headers, body) = parse_response(raw);
assert_eq!(status, 201);
assert_eq!(body, b"{\"ok\":true}");
assert!(headers.contains(&("content-type".to_string(), "application/json".to_string())));
assert!(
!headers.iter().any(|(n, _)| n == "connection"),
"{headers:?}"
);
}
#[test]
fn a_malformed_response_is_reported_as_a_bad_gateway() {
let (status, _, _) = parse_response(b"garbage");
assert_eq!(status, 502);
}
}