use axum::{
extract::{
FromRequestParts, State,
ws::{Message as AxumMessage, WebSocket, WebSocketUpgrade},
},
http::HeaderMap,
response::{IntoResponse, Response},
};
use bytes::Bytes;
use futures::{Sink, Stream};
use std::{
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use tokio::time::timeout;
use tracing::{debug, info, trace, warn};
use iroh_relay::http::ProtocolVersion;
use iroh_relay::{
ExportKeyingMaterial, KeyCache,
protos::{handshake, streams::StreamError},
server::{
Access, AccessControl, AllowAll, ClientRequest, DynAccessControl,
Metrics,
client::Config,
clients::Clients,
streams::{Bucket, RelayedStream},
},
};
use std::future::Future;
use crate::RelayAllowlist;
#[derive(Clone, Copy, Debug)]
pub struct RelayClientRxRateLimit {
pub bytes_per_second: std::num::NonZeroU32,
pub burst_bytes: std::num::NonZeroU32,
}
#[derive(Clone, Debug)]
pub struct RelayState {
pub key_cache: KeyCache,
pub access: Arc<dyn DynAccessControl>,
pub metrics: Arc<Metrics>,
pub write_timeout: std::time::Duration,
pub clients: Clients,
pub rate_limit: Option<RelayClientRxRateLimit>,
}
impl RelayState {
pub fn new(
key_cache: KeyCache,
access: Arc<dyn DynAccessControl>,
metrics: Arc<Metrics>,
rate_limit: Option<RelayClientRxRateLimit>,
) -> Self {
Self {
key_cache,
access,
metrics,
write_timeout: iroh_relay::defaults::timeouts::SERVER_WRITE_TIMEOUT,
clients: Clients::default(),
rate_limit,
}
}
}
pub async fn relay_handler(
State(state): State<RelayState>,
request: axum::extract::Request,
) -> Response {
let (mut parts, _body) = request.into_parts();
let client_auth_header = parts
.headers
.get(iroh_relay::http::CLIENT_AUTH_HEADER)
.cloned();
let has_auth = client_auth_header.is_some();
debug!(?has_auth, "Relay WebSocket upgrade request");
let ws =
match WebSocketUpgrade::from_request_parts(&mut parts, &state).await {
Ok(ws) => ws,
Err(rejection) => return rejection.into_response(),
};
let client_request_parts = snapshot_request_parts(&parts);
let ws = ws.protocols([ProtocolVersion::V1.to_str()]);
ws.on_upgrade(move |socket| async move {
if let Err(e) = handle_relay_websocket(
socket,
state,
client_auth_header,
client_request_parts,
)
.await
{
let msg = e.to_string().to_lowercase();
let is_expected = msg.contains("timed out")
|| msg.contains("connection reset")
|| msg.contains("connection closed")
|| msg.contains("broken pipe")
|| msg.contains("not allowed")
|| msg.contains("denied")
|| msg.contains("websocket protocol error");
if is_expected {
debug!("Relay WebSocket ended: {msg}");
} else {
warn!("Error handling relay WebSocket: {msg}");
}
}
})
}
struct RateLimitState {
bucket: Bucket,
sleep: Option<Pin<Box<tokio::time::Sleep>>>,
}
impl RateLimitState {
fn new(limit: RelayClientRxRateLimit) -> Self {
let bucket = Bucket::new(
limit.burst_bytes.get() as i64,
limit.bytes_per_second.get() as i64,
std::time::Duration::from_millis(100),
)
.expect(
"RelayClientRxRateLimit fields are NonZeroU32 so Bucket::new \
cannot fail with InvalidBucketConfig",
);
Self {
bucket,
sleep: None,
}
}
fn poll_throttle(&mut self, cx: &mut Context<'_>) -> Poll<()> {
if let Some(sleep) = self.sleep.as_mut() {
match sleep.as_mut().poll(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(()) => self.sleep = None,
}
}
Poll::Ready(())
}
fn charge(&mut self, bytes: usize) {
if let Err(deadline) = self.bucket.consume(bytes) {
self.sleep = Some(Box::pin(tokio::time::sleep_until(deadline)));
}
}
}
struct AxumWebSocketAdapter {
inner: Pin<Box<WebSocket>>,
rate_limit: Option<RateLimitState>,
}
impl AxumWebSocketAdapter {
fn new(
socket: WebSocket,
rate_limit: Option<RelayClientRxRateLimit>,
) -> Self {
Self {
inner: Box::pin(socket),
rate_limit: rate_limit.map(RateLimitState::new),
}
}
}
impl Stream for AxumWebSocketAdapter {
type Item = Result<Bytes, StreamError>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
if let Some(rl) = self.rate_limit.as_mut()
&& rl.poll_throttle(cx).is_pending()
{
return Poll::Pending;
}
loop {
match self.inner.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok(msg))) => match msg {
AxumMessage::Binary(data) => {
if let Some(rl) = self.rate_limit.as_mut() {
rl.charge(data.len());
}
return Poll::Ready(Some(Ok(data)));
}
AxumMessage::Close(_) => return Poll::Ready(None),
_ => continue,
},
Poll::Ready(Some(Err(e))) => {
return Poll::Ready(Some(Err(StreamError::from_std(e))));
}
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
}
}
}
}
impl Sink<Bytes> for AxumWebSocketAdapter {
type Error = StreamError;
fn poll_ready(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
match self.inner.as_mut().poll_ready(cx) {
Poll::Ready(Ok(())) => Poll::Ready(Ok(())),
Poll::Ready(Err(e)) => Poll::Ready(Err(StreamError::from_std(e))),
Poll::Pending => Poll::Pending,
}
}
fn start_send(
mut self: Pin<&mut Self>,
item: Bytes,
) -> Result<(), Self::Error> {
self.inner
.as_mut()
.start_send(AxumMessage::Binary(item))
.map_err(StreamError::from_std)
}
fn poll_flush(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
match self.inner.as_mut().poll_flush(cx) {
Poll::Ready(Ok(())) => Poll::Ready(Ok(())),
Poll::Ready(Err(e)) => Poll::Ready(Err(StreamError::from_std(e))),
Poll::Pending => Poll::Pending,
}
}
fn poll_close(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
match self.inner.as_mut().poll_close(cx) {
Poll::Ready(Ok(())) => Poll::Ready(Ok(())),
Poll::Ready(Err(e)) => Poll::Ready(Err(StreamError::from_std(e))),
Poll::Pending => Poll::Pending,
}
}
}
impl ExportKeyingMaterial for AxumWebSocketAdapter {
fn export_keying_material<T: AsMut<[u8]>>(
&self,
_output: T,
_label: &[u8],
_context: Option<&[u8]>,
) -> Option<T> {
None
}
}
impl Unpin for AxumWebSocketAdapter {}
async fn handle_relay_websocket(
socket: WebSocket,
state: RelayState,
client_auth_header: Option<http::HeaderValue>,
request_parts: http::request::Parts,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
trace!("Relay WebSocket connection established");
let mut adapter = AxumWebSocketAdapter::new(socket, state.rate_limit);
const HANDSHAKE_TIMEOUT: std::time::Duration =
std::time::Duration::from_secs(30);
let (endpoint_id, guard) = timeout(HANDSHAKE_TIMEOUT, async {
let authentication =
handshake::serverside(&mut adapter, client_auth_header).await?;
debug!(?authentication.mechanism, "accept: verified authentication");
let client_request = ClientRequest::new(
authentication.client_key,
ProtocolVersion::V1,
request_parts,
);
let guard = authentication
.authorize_with(&client_request, &state.access, &mut adapter)
.await?;
let endpoint_id = guard.endpoint_id();
Ok::<_, Box<dyn std::error::Error + Send + Sync>>((endpoint_id, guard))
})
.await
.map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::TimedOut,
"Relay handshake timed out",
)
})??;
debug!(key = %endpoint_id.fmt_short(), "accept: verified authorization");
let io = RelayedStream::new(adapter, state.key_cache.clone());
trace!("accept: build client conn");
let mut client_conn_builder = Config::new(guard, io, ProtocolVersion::V1);
client_conn_builder.write_timeout = state.write_timeout;
state
.clients
.register(client_conn_builder, state.metrics.clone());
info!(key = %endpoint_id.fmt_short(), "relay client registered");
Ok(())
}
fn snapshot_request_parts(
parts: &http::request::Parts,
) -> http::request::Parts {
let mut new_parts = http::Request::builder()
.method(parts.method.clone())
.uri(parts.uri.clone())
.body(())
.expect("building an empty http::Request must succeed")
.into_parts()
.0;
new_parts.headers = parts.headers.clone();
new_parts.version = parts.version;
new_parts
}
pub async fn relay_probe_handler() -> (
axum::http::StatusCode,
axum::response::AppendHeaders<[(axum::http::HeaderName, &'static str); 1]>,
) {
(
axum::http::StatusCode::OK,
axum::response::AppendHeaders([(
axum::http::header::ACCESS_CONTROL_ALLOW_ORIGIN,
"*",
)]),
)
}
pub async fn captive_portal_handler(
headers: HeaderMap,
) -> axum::response::Response {
let mut response = axum::response::Response::builder()
.status(axum::http::StatusCode::NO_CONTENT);
if let Some(challenge) = headers.get("X-Iroh-Challenge")
&& let Ok(challenge_str) = challenge.to_str()
&& !challenge_str.is_empty()
&& challenge_str.len() < 64
&& challenge_str
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_')
{
response = response
.header("X-Iroh-Response", format!("response {challenge_str}"));
}
response
.body(axum::body::Body::empty())
.unwrap_or_else(|_| {
axum::response::Response::builder()
.status(axum::http::StatusCode::NO_CONTENT)
.body(axum::body::Body::empty())
.unwrap()
})
}
pub fn create_relay_state(
rate_limit: Option<RelayClientRxRateLimit>,
) -> RelayState {
let key_cache = KeyCache::new(1024);
let access = Arc::new(AllowAll);
let metrics = Arc::new(Metrics::default());
RelayState::new(key_cache, access, metrics, rate_limit)
}
struct AllowlistAccess {
allowlist: RelayAllowlist,
}
impl AllowlistAccess {
fn new(allowlist: RelayAllowlist) -> Self {
Self { allowlist }
}
}
impl std::fmt::Debug for AllowlistAccess {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AllowlistAccess").finish()
}
}
impl AccessControl for AllowlistAccess {
async fn on_connect(&self, request: &ClientRequest) -> Access {
let endpoint_id = request.endpoint_id();
let allowed = self.allowlist.is_allowed(&endpoint_id);
tracing::debug!(
key = %endpoint_id.fmt_short(),
allowed,
allowlist_size = self.allowlist.len(),
"Relay access check"
);
if allowed {
Access::Allow
} else {
Access::Deny { reason: None }
}
}
}
pub fn create_relay_state_with_allowlist(
allowlist: RelayAllowlist,
rate_limit: Option<RelayClientRxRateLimit>,
) -> RelayState {
let key_cache = KeyCache::new(1024);
let access = Arc::new(AllowlistAccess::new(allowlist));
let metrics = Arc::new(Metrics::default());
RelayState::new(key_cache, access, metrics, rate_limit)
}
#[cfg(test)]
mod tests {
use super::*;
use axum::Router;
use axum::routing::get;
use futures::{SinkExt, StreamExt};
use iroh_base::{RelayUrl, SecretKey};
use std::net::Ipv4Addr;
use tokio::net::TcpListener;
use tracing::{info, instrument};
use iroh_dns::dns::DnsResolver;
use iroh_relay::{
client::ClientBuilder,
protos::relay::{ClientToRelayMsg, Datagrams, RelayToClientMsg},
tls::CaTlsConfig,
};
#[tokio::test]
#[instrument]
async fn axum_relay_integration() -> Result<(), Box<dyn std::error::Error>>
{
let state = create_relay_state(None);
let app = Router::new()
.route("/relay", get(relay_handler))
.with_state(state.clone());
let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await?;
let addr = listener.local_addr()?;
info!("Axum relay server listening on {}", addr);
let server_handle = tokio::spawn(async move {
axum::serve(listener, app).await.expect("server error");
});
let relay_url = format!("http://{}/relay", addr);
let relay_url: RelayUrl = relay_url.parse()?;
let tls_client_config = CaTlsConfig::default()
.client_config(iroh_relay::tls::default_provider())?;
let a_secret_key = SecretKey::generate();
let a_key = a_secret_key.public();
let resolver = DnsResolver::new();
info!("Connecting client A");
let mut client_a = ClientBuilder::new(
relay_url.clone(),
a_secret_key,
resolver.clone(),
)
.tls_client_config(tls_client_config.clone())
.connect()
.await?;
let b_secret_key = SecretKey::generate();
let b_key = b_secret_key.public();
info!("Connecting client B");
let mut client_b = ClientBuilder::new(
relay_url.clone(),
b_secret_key,
resolver.clone(),
)
.tls_client_config(tls_client_config)
.connect()
.await?;
info!("Sending message from A to B");
let msg = Datagrams::from("hello from A");
client_a
.send(ClientToRelayMsg::Datagrams {
dst_endpoint_id: b_key,
datagrams: msg.clone(),
})
.await?;
let received = tokio::time::timeout(
std::time::Duration::from_secs(2),
client_b.next(),
)
.await
.expect("timeout waiting for message")
.expect("stream ended")?;
match received {
RelayToClientMsg::Datagrams {
remote_endpoint_id,
datagrams,
} => {
assert_eq!(remote_endpoint_id, a_key, "Wrong sender");
assert_eq!(datagrams, msg, "Message content mismatch");
info!("Successfully received message on client B");
}
other => panic!("Unexpected message type: {:?}", other),
}
info!("Sending message from B to A");
let msg2 = Datagrams::from("hello from B");
client_b
.send(ClientToRelayMsg::Datagrams {
dst_endpoint_id: a_key,
datagrams: msg2.clone(),
})
.await?;
let received = tokio::time::timeout(
std::time::Duration::from_secs(2),
client_a.next(),
)
.await
.expect("timeout waiting for message")
.expect("stream ended")?;
match received {
RelayToClientMsg::Datagrams {
remote_endpoint_id,
datagrams,
} => {
assert_eq!(remote_endpoint_id, b_key, "Wrong sender");
assert_eq!(datagrams, msg2, "Message content mismatch");
info!("Successfully received message on client A");
}
other => panic!("Unexpected message type: {:?}", other),
}
drop(client_a);
drop(client_b);
server_handle.abort();
info!("Test completed successfully");
Ok(())
}
#[tokio::test]
#[instrument]
async fn axum_relay_rate_limited() -> Result<(), Box<dyn std::error::Error>>
{
use std::num::NonZeroU32;
use std::time::Instant as StdInstant;
let rate_limit = RelayClientRxRateLimit {
bytes_per_second: NonZeroU32::new(8 * 1024).unwrap(),
burst_bytes: NonZeroU32::new(8 * 1024).unwrap(),
};
let state = create_relay_state(Some(rate_limit));
let app = Router::new()
.route("/relay", get(relay_handler))
.with_state(state.clone());
let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await?;
let addr = listener.local_addr()?;
let server_handle = tokio::spawn(async move {
axum::serve(listener, app).await.expect("server error");
});
let relay_url = format!("http://{}/relay", addr);
let relay_url: RelayUrl = relay_url.parse()?;
let tls_client_config = CaTlsConfig::default()
.client_config(iroh_relay::tls::default_provider())?;
let a_secret_key = SecretKey::generate();
let resolver = DnsResolver::new();
let mut client_a = ClientBuilder::new(
relay_url.clone(),
a_secret_key,
resolver.clone(),
)
.tls_client_config(tls_client_config.clone())
.connect()
.await?;
let b_secret_key = SecretKey::generate();
let b_key = b_secret_key.public();
let mut client_b = ClientBuilder::new(
relay_url.clone(),
b_secret_key,
resolver.clone(),
)
.tls_client_config(tls_client_config)
.connect()
.await?;
let payload = vec![0u8; 1024];
let datagrams = Datagrams::from(payload);
let started = StdInstant::now();
for _ in 0..24 {
client_a
.send(ClientToRelayMsg::Datagrams {
dst_endpoint_id: b_key,
datagrams: datagrams.clone(),
})
.await?;
}
let mut received = 0usize;
while received < 24 {
let msg = tokio::time::timeout(
std::time::Duration::from_secs(10),
client_b.next(),
)
.await
.expect("timeout waiting for message")
.expect("stream ended")?;
if matches!(msg, RelayToClientMsg::Datagrams { .. }) {
received += 1;
}
}
let elapsed = started.elapsed();
assert!(
elapsed >= std::time::Duration::from_millis(1500),
"rate-limit appears not to be applied: elapsed = {elapsed:?}"
);
assert!(
elapsed < std::time::Duration::from_secs(9),
"rate-limit appears stuck: elapsed = {elapsed:?}"
);
drop(client_a);
drop(client_b);
server_handle.abort();
Ok(())
}
use std::num::NonZeroU32;
fn poll_throttle_now(state: &mut RateLimitState) -> Poll<()> {
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
state.poll_throttle(&mut cx)
}
#[tokio::test(start_paused = true)]
async fn rate_limit_passes_under_burst() {
let limit = RelayClientRxRateLimit {
bytes_per_second: NonZeroU32::new(1024).unwrap(),
burst_bytes: NonZeroU32::new(1024).unwrap(),
};
let mut state = RateLimitState::new(limit);
state.charge(512);
assert!(matches!(poll_throttle_now(&mut state), Poll::Ready(())));
assert!(state.sleep.is_none());
}
#[tokio::test(start_paused = true)]
async fn rate_limit_paces_overrun() {
let limit = RelayClientRxRateLimit {
bytes_per_second: NonZeroU32::new(1024).unwrap(),
burst_bytes: NonZeroU32::new(1024).unwrap(),
};
let mut state = RateLimitState::new(limit);
state.charge(512);
assert!(matches!(poll_throttle_now(&mut state), Poll::Ready(())));
assert!(state.sleep.is_none());
state.charge(1024);
assert!(state.sleep.is_some());
assert!(matches!(poll_throttle_now(&mut state), Poll::Pending));
let started = tokio::time::Instant::now();
std::future::poll_fn(|cx| state.poll_throttle(cx)).await;
let elapsed = started.elapsed();
assert!(state.sleep.is_none());
assert!(
elapsed >= std::time::Duration::from_millis(100),
"throttle cleared too soon: elapsed = {elapsed:?}"
);
}
#[tokio::test]
async fn rate_limit_state_construction_safe_for_low_rates() {
let limit = RelayClientRxRateLimit {
bytes_per_second: NonZeroU32::new(100).unwrap(),
burst_bytes: NonZeroU32::new(10).unwrap(),
};
RateLimitState::new(limit);
}
#[tokio::test]
#[should_panic]
async fn rate_limit_state_construction_panics_for_zero_burst() {
let limit = RelayClientRxRateLimit {
bytes_per_second: NonZeroU32::new(100).unwrap(),
burst_bytes: NonZeroU32::new(0).unwrap(),
};
RateLimitState::new(limit);
}
}