use crate::network::{
read_len_prefix, read_len_prefix_with_slack, relay_credential, write_encoded, Message,
MessageSink, MessageSource, MessageTransport, NetworkConfig, ServerIdentity, SessionToken,
};
use anyhow::{Context, Result};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::mpsc;
use tokio::sync::Mutex;
use tokio_rustls::client::TlsStream;
use tracing::{info, warn};
#[derive(Debug, Serialize, Deserialize)]
pub struct RelayRegister {
pub session: String,
pub role: RelayRole,
pub token: String,
pub protocol_version: u8,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum RelayRole {
Host,
Viewer,
}
const MAX_REGISTRATION_BYTES: usize = 4096;
const MAX_PEER_QUEUE_BYTES: usize = 8 * 1024 * 1024;
const MAX_PEER_QUEUE_MESSAGES: usize = 64;
const MAX_SESSIONS: usize = 1024;
fn same_credential(expected: &str, presented: &str) -> bool {
let a = expected.as_bytes();
let b = presented.as_bytes();
let mut diff = (a.len() ^ b.len()) as u8;
for i in 0..a.len().max(b.len()) {
diff |= a.get(i).copied().unwrap_or(0) ^ b.get(i).copied().unwrap_or(0);
}
diff == 0
}
const MAX_VIEWERS_PER_SESSION: usize = 64;
const AUTH_FAILURES_ALLOWED: u32 = 8;
const AUTH_FAILURE_WINDOW: Duration = Duration::from_secs(60);
const SESSION_IDLE_TTL: Duration = Duration::from_secs(3600);
pub fn generate_session_code() -> String {
use rand::Rng;
const CHARS: &[u8] = b"ABCDEFGHJKLMNPQRSTUVWXYZ23456789"; let mut rng = rand::thread_rng();
(0..6)
.map(|_| CHARS[rng.gen_range(0..CHARS.len())] as char)
.collect()
}
#[derive(Clone)]
struct Peer {
gen: u64,
id: u32,
tx: mpsc::Sender<Vec<u8>>,
}
pub(crate) struct Session {
host: Option<Peer>,
viewers: Vec<Peer>,
next_viewer_id: u32,
last_seen: Instant,
}
pub(crate) type Sessions = Arc<Mutex<HashMap<String, Session>>>;
pub struct RelayTransport {
stream: TlsStream<TcpStream>,
}
impl RelayTransport {
pub async fn connect(
relay_addr: std::net::SocketAddr,
server_name: &str,
pin: &[u8],
session: String,
token: SessionToken,
role: RelayRole,
) -> Result<Self> {
let config = NetworkConfig::client_tls_config(pin)?;
let server_name = rustls::pki_types::ServerName::try_from(server_name.to_owned())
.map_err(|e| anyhow::anyhow!("Invalid relay server name '{server_name}': {e}"))?;
let connector = tokio_rustls::TlsConnector::from(Arc::new(config));
let tcp = TcpStream::connect(relay_addr)
.await
.with_context(|| format!("Failed to connect to the relay at {relay_addr}"))?;
let mut stream = connector
.connect(server_name, tcp)
.await
.with_context(|| format!("TLS handshake with the relay at {relay_addr} failed"))?;
let reg = RelayRegister {
session,
role,
token: relay_credential(&token),
protocol_version: crate::network::PROTOCOL_VERSION,
};
let bytes = bincode::serialize(®)?;
stream
.write_all(&(bytes.len() as u32).to_le_bytes())
.await?;
stream.write_all(&bytes).await?;
stream.flush().await?;
let mut ack = [0u8; 1];
tokio::time::timeout(Duration::from_secs(10), stream.read_exact(&mut ack))
.await
.context("Timed out waiting for the relay to accept the registration")?
.context("The relay closed the connection during registration")?;
if ack[0] == 0 {
anyhow::bail!(
"no host in this session yet — is the sharer still running? (the relay accepted the registration, but the session is empty)"
);
}
if ack[0] != crate::network::PROTOCOL_VERSION {
anyhow::bail!(
"The relay rejected this viewer: check --session and --token \
(relay speaks protocol {}, offered {})",
ack[0],
crate::network::PROTOCOL_VERSION
);
}
Ok(Self { stream })
}
}
pub const RELAY_ID_BYTES: usize = 4;
pub const RELAY_CONTROL_ID: u32 = 0xFFFF_FFFF;
const K_CONTROL_VIEWER_HERE: u8 = 0xEF;
const K_CONTROL_VIEWER_LEFT: u8 = 0xEE;
pub(crate) const K_CONTROL_BROWSER_HERE: u8 = 0xED;
pub struct RelayFan {
pub rx: mpsc::Receiver<(u32, Vec<u8>)>,
pub tx: mpsc::Sender<(u32, Vec<u8>)>,
}
pub async fn read_host_frame<R: tokio::io::AsyncRead + Unpin>(
reader: &mut R,
) -> Result<Option<(u32, Vec<u8>)>> {
let mut len_buf = [0u8; 4];
match reader.read_exact(&mut len_buf).await {
Ok(_) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
Err(e) => return Err(e.into()),
}
let len = read_len_prefix_with_slack(&len_buf, RELAY_ID_BYTES)?;
if len < RELAY_ID_BYTES {
anyhow::bail!("host-leg frame too short for a peer id: {len} bytes");
}
let mut rest = vec![0u8; len];
reader.read_exact(&mut rest).await?;
let id = u32::from_le_bytes(rest[0..4].try_into().unwrap());
Ok(Some((id, rest[4..].to_vec())))
}
const FAN_SESSION_QUEUE: usize = 64;
fn tag_host_frame(id: u32, envelope: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(8 + envelope.len());
out.extend_from_slice(&((envelope.len() + 4) as u32).to_le_bytes());
out.extend_from_slice(&id.to_le_bytes());
out.extend_from_slice(envelope);
out
}
struct RelaySink {
stream: tokio::io::WriteHalf<TlsStream<TcpStream>>,
}
#[async_trait]
impl MessageSink for RelaySink {
async fn send(&mut self, msg: &Message) -> Result<()> {
self.send_encoded(&msg.encode()?).await
}
async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()> {
write_encoded(&mut self.stream, bytes).await
}
}
struct RelaySource {
stream: tokio::io::ReadHalf<TlsStream<TcpStream>>,
}
#[async_trait]
impl MessageSource for RelaySource {
async fn recv(&mut self) -> Result<Message> {
Message::read_framed(&mut self.stream).await
}
async fn recv_raw(&mut self) -> Result<Vec<u8>> {
Message::read_envelope(&mut self.stream).await
}
}
impl RelayTransport {
pub fn into_fan(self: Box<Self>) -> RelayFan {
let (fan_tx_in, fan_rx_in) = mpsc::channel::<(u32, Vec<u8>)>(FAN_SESSION_QUEUE);
let (fan_tx_out, mut fan_rx_out) = mpsc::channel::<(u32, Vec<u8>)>(FAN_SESSION_QUEUE);
let (mut read_half, mut write_half) = tokio::io::split(self.stream);
tokio::spawn(async move {
loop {
match read_host_frame(&mut read_half).await {
Ok(Some((id, env))) => {
if fan_tx_in.send((id, env)).await.is_err() {
break;
}
}
Ok(None) => break,
Err(_) => break,
}
}
});
tokio::spawn(async move {
while let Some((id, env)) = fan_rx_out.recv().await {
if write_half
.write_all(&tag_host_frame(id, &env))
.await
.is_err()
{
break;
}
}
});
RelayFan {
rx: fan_rx_in,
tx: fan_tx_out,
}
}
}
#[async_trait]
impl MessageTransport for RelayTransport {
async fn send(&mut self, msg: &Message) -> Result<()> {
write_encoded(&mut self.stream, &msg.encode()?).await
}
async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()> {
write_encoded(&mut self.stream, bytes).await
}
async fn recv(&mut self) -> Result<Message> {
Message::read_framed(&mut self.stream).await
}
fn split(self: Box<Self>) -> (Box<dyn MessageSink>, Box<dyn MessageSource>) {
let (read, write) = tokio::io::split(self.stream);
(
Box::new(RelaySink { stream: write }),
Box::new(RelaySource { stream: read }),
)
}
}
pub(crate) struct BrowserViewer {
pub(crate) id: u32,
pub(crate) rx: mpsc::Receiver<Vec<u8>>,
gen: u64,
}
pub(crate) async fn register_browser_viewer(
sessions: &Sessions,
next_gen: &AtomicU64,
session_id: &str,
) -> Result<BrowserViewer> {
let (tx, rx) = mpsc::channel::<Vec<u8>>(MAX_PEER_QUEUE_MESSAGES);
let mut map = sessions.lock().await;
let session = map
.get_mut(session_id)
.with_context(|| format!("no session '{session_id}'"))?;
if session.viewers.len() >= MAX_VIEWERS_PER_SESSION {
anyhow::bail!(
"session '{session_id}' already has {} viewers (max {MAX_VIEWERS_PER_SESSION})",
session.viewers.len()
);
}
let host = session
.host
.as_ref()
.with_context(|| format!("session '{session_id}' has no host"))?;
let id = session.next_viewer_id;
session.next_viewer_id = session.next_viewer_id.saturating_add(1).max(1);
if id == RELAY_CONTROL_ID {
anyhow::bail!("session '{session_id}' exhausted its viewer ids");
}
let gen = next_gen.fetch_add(1, Ordering::Relaxed);
session.viewers.push(Peer { gen, id, tx });
session.last_seen = Instant::now();
let mut frame = Vec::with_capacity(9);
frame.extend_from_slice(&9u32.to_le_bytes());
frame.extend_from_slice(&RELAY_CONTROL_ID.to_le_bytes());
frame.push(K_CONTROL_BROWSER_HERE);
frame.extend_from_slice(&id.to_le_bytes());
let _ = host.tx.try_send(frame);
Ok(BrowserViewer { id, rx, gen })
}
pub(crate) async fn unregister_browser_viewer(
sessions: &Sessions,
session_id: &str,
viewer: &BrowserViewer,
) {
let mut map = sessions.lock().await;
let Some(session) = map.get_mut(session_id) else {
return;
};
session.viewers.retain(|v| v.gen != viewer.gen);
if let Some(host) = session.host.as_ref() {
let mut frame = Vec::with_capacity(9);
frame.extend_from_slice(&9u32.to_le_bytes());
frame.extend_from_slice(&RELAY_CONTROL_ID.to_le_bytes());
frame.push(K_CONTROL_VIEWER_LEFT);
frame.extend_from_slice(&viewer.id.to_le_bytes());
let _ = host.tx.try_send(frame);
}
if session.host.is_none() && session.viewers.is_empty() {
map.remove(session_id);
}
}
pub(crate) async fn forward_browser_payload(
sessions: &Sessions,
session_id: &str,
viewer_id: u32,
payload: &[u8],
) -> Result<()> {
let target = {
let mut map = sessions.lock().await;
map.get_mut(session_id).and_then(|s| {
s.last_seen = Instant::now();
s.host.as_ref().map(|h| h.tx.clone())
})
};
let Some(host) = target else {
anyhow::bail!("session '{session_id}' has no host");
};
let mut frame = Vec::with_capacity(8 + payload.len());
frame.extend_from_slice(&((payload.len() + RELAY_ID_BYTES) as u32).to_le_bytes());
frame.extend_from_slice(&viewer_id.to_le_bytes());
frame.extend_from_slice(payload);
host.try_send(frame)
.map_err(|_| anyhow::anyhow!("host is not draining; dropping the browser viewer"))
}
pub(crate) async fn session_has_host(sessions: &Sessions, session_id: &str) -> bool {
sessions
.lock()
.await
.get(session_id)
.is_some_and(|s| s.host.is_some())
}
pub(crate) type SessionTable = Sessions;
pub async fn run_relay_server(
listener: TcpListener,
identity: Arc<ServerIdentity>,
token: SessionToken,
) -> Result<()> {
run_relay_server_with(listener, identity, token, false).await
}
pub async fn run_relay_server_with(
listener: TcpListener,
identity: Arc<ServerIdentity>,
token: SessionToken,
serve_web: bool,
) -> Result<()> {
let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(NetworkConfig::relay_server_config(
&identity, serve_web,
)?));
info!("Relay server listening on {} (TLS)", listener.local_addr()?);
println!(
"Relay certificate fingerprint (sha256): {}",
identity.fingerprint
);
let sessions: Sessions = Arc::new(Mutex::new(HashMap::new()));
let next_gen = Arc::new(AtomicU64::new(1));
let failures: Arc<Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>> =
Arc::new(Mutex::new(HashMap::new()));
{
let sessions = sessions.clone();
let failures = failures.clone();
tokio::spawn(async move { sweep(sessions, failures).await });
}
let ctx = RelayCtx {
acceptor,
sessions,
next_gen,
token,
failures,
serve_web,
};
loop {
let (tcp, peer) = listener.accept().await?;
let ctx = ctx.clone();
tokio::spawn(async move {
if let Err(e) = handle_client(tcp, peer, ctx).await {
warn!("Relay peer {peer}: {e}");
}
});
}
}
#[derive(Clone)]
struct RelayCtx {
acceptor: tokio_rustls::TlsAcceptor,
sessions: Sessions,
next_gen: Arc<AtomicU64>,
token: SessionToken,
failures: Arc<tokio::sync::Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>>,
serve_web: bool,
}
async fn sweep(
sessions: Sessions,
failures: Arc<tokio::sync::Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>>,
) {
let mut ticker = tokio::time::interval(Duration::from_secs(60));
loop {
ticker.tick().await;
{
let mut map = sessions.lock().await;
map.retain(|_, s| s.last_seen.elapsed() < SESSION_IDLE_TTL);
}
{
let mut map = failures.lock().await;
map.retain(|_, (_, at)| at.elapsed() < AUTH_FAILURE_WINDOW);
}
}
}
async fn handle_client(tcp: TcpStream, peer: std::net::SocketAddr, ctx: RelayCtx) -> Result<()> {
let stream = match ctx.acceptor.accept(tcp).await {
Ok(s) => s,
Err(e) => {
tracing::debug!("Relay peer {peer}: TLS handshake failed (bare connect?): {e}");
return Ok(());
}
};
if ctx.serve_web && stream.get_ref().1.alpn_protocol() == Some(b"http/1.1") {
return crate::relay_web::handle_http(
stream,
&ctx.sessions,
ctx.next_gen.as_ref(),
&ctx.token,
)
.await;
}
if record_and_check_rate(&ctx.failures, peer).await {
anyhow::bail!(
"too many failed registrations (max {AUTH_FAILURES_ALLOWED} per {}s)",
AUTH_FAILURE_WINDOW.as_secs()
);
}
handle_relay_client(stream, peer, ctx).await
}
async fn handle_relay_client(
stream: tokio_rustls::server::TlsStream<TcpStream>,
peer: std::net::SocketAddr,
ctx: RelayCtx,
) -> Result<()> {
let RelayCtx {
sessions,
next_gen,
token: expected_token,
failures,
..
} = ctx;
let (mut read_half, mut write_half) = tokio::io::split(stream);
let reg = read_registration(&mut read_half).await?;
if !same_credential(relay_credential(&expected_token).as_str(), ®.token) {
if reg.protocol_version == 0 || reg.protocol_version != crate::network::PROTOCOL_VERSION {
let client = if reg.protocol_version == 0 {
"pre-0.1.5 (no version in registration)".to_string()
} else {
format!("protocol {}", reg.protocol_version)
};
warn!(
"Relay: rejected {:?} from {peer}: client speaks {client}, this relay speaks protocol {} — upgrade the client, the token is not the problem",
reg.role,
crate::network::PROTOCOL_VERSION
);
anyhow::bail!(
"client speaks {client}, this relay speaks protocol {}: upgrade the client",
crate::network::PROTOCOL_VERSION
);
}
warn!(
"Relay: rejected {:?} from {peer} (bad credential)",
reg.role
);
anyhow::bail!("bad {:?} credential", reg.role);
}
failures.lock().await.remove(&peer);
if reg.role == RelayRole::Viewer {
let empty = {
let map = sessions.lock().await;
map.get(®.session)
.map(|s| s.host.is_none())
.unwrap_or(true)
};
if empty {
info!(
"Relay: viewer from {peer} parked in session '{}' (no host yet)",
reg.session
);
}
}
write_half
.write_all(&[crate::network::PROTOCOL_VERSION])
.await?;
write_half.flush().await?;
info!(
"Relay: {:?} joined session '{}' from {peer}",
reg.role, reg.session
);
let host_channel = reg.role == RelayRole::Host;
let (tx, mut rx) = mpsc::channel::<Vec<u8>>(if host_channel {
MAX_PEER_QUEUE_MESSAGES + MAX_VIEWERS_PER_SESSION
} else {
MAX_PEER_QUEUE_MESSAGES
});
let gen = next_gen.fetch_add(1, Ordering::Relaxed);
let mut my_id: u32 = 0;
let mut here_replay: Vec<u32> = Vec::new();
let peer_entry = {
let mut map = sessions.lock().await;
let session = map.entry(reg.session.clone()).or_insert_with(|| Session {
host: None,
viewers: Vec::new(),
next_viewer_id: 1,
last_seen: Instant::now(),
});
session.last_seen = Instant::now();
let entry = match reg.role {
RelayRole::Host => {
let entry = Peer {
gen,
id: 0,
tx: tx.clone(),
};
if let Some(previous) = session.host.replace(entry.clone()) {
drop(previous);
}
here_replay = session.viewers.iter().map(|v| v.id).collect();
entry
}
RelayRole::Viewer => {
if session.viewers.len() >= MAX_VIEWERS_PER_SESSION {
anyhow::bail!(
"session '{}' already has {} viewers (max {MAX_VIEWERS_PER_SESSION})",
reg.session,
session.viewers.len()
);
}
let id = session.next_viewer_id;
session.next_viewer_id = session.next_viewer_id.saturating_add(1).max(1);
if id == RELAY_CONTROL_ID {
anyhow::bail!("session '{}' exhausted its viewer ids", reg.session);
}
my_id = id;
let entry = Peer {
gen,
id,
tx: tx.clone(),
};
session.viewers.push(entry.clone());
entry
}
};
if map.len() > MAX_SESSIONS {
anyhow::bail!(
"relay is hosting {} sessions (max {MAX_SESSIONS})",
map.len()
);
}
entry
};
drop(tx);
let peer_role = reg.role;
let session_id = reg.session.clone();
let mut peer_entry = Some(peer_entry);
let mut queued_bytes = 0usize;
for id in here_replay {
let mut frame = Vec::with_capacity(9);
frame.extend_from_slice(&9u32.to_le_bytes());
frame.extend_from_slice(&RELAY_CONTROL_ID.to_le_bytes());
frame.push(K_CONTROL_VIEWER_HERE);
frame.extend_from_slice(&id.to_le_bytes());
if write_half.write_all(&frame).await.is_err() {
break;
}
}
loop {
tokio::select! {
inbound = async {
if host_channel {
read_host_frame(&mut read_half).await.map(|o| {
o.map(|(id, env)| {
let mut framed = Vec::with_capacity(8 + env.len());
framed.extend_from_slice(&id.to_le_bytes());
framed.extend_from_slice(&env);
framed
})
})
} else {
read_frame(&mut read_half).await
}
} => {
let Some(framed) = inbound? else { break }; let len = framed.len();
if host_channel {
let id = u32::from_le_bytes(framed[0..4].try_into().unwrap());
if id == RELAY_CONTROL_ID {
warn!("Relay: host sent a control frame; dropping it");
continue;
}
let target = {
let map = sessions.lock().await;
map.get(&session_id).and_then(|session| {
session.viewers.iter().find(|v| v.id == id).map(|v| v.tx.clone())
})
};
let Some(target) = target else {
warn!("Relay: frame for unknown viewer id {id}; dropping it");
continue;
};
let env = framed.get(4..).unwrap_or(&[]);
let mut out = Vec::with_capacity(4 + env.len());
out.extend_from_slice(&(env.len() as u32).to_le_bytes());
out.extend_from_slice(env);
match target.try_send(out) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
warn!("Relay: dropping host: viewer {id} not draining");
break;
}
Err(mpsc::error::TrySendError::Closed(_)) => {}
}
if let Some(session) = sessions.lock().await.get_mut(&session_id) {
session.last_seen = Instant::now();
}
} else {
let framed = {
let env = framed.get(4..).unwrap_or(&[]);
let mut tagged = Vec::with_capacity(8 + env.len());
tagged.extend_from_slice(
&((env.len() + 4) as u32).to_le_bytes(),
);
tagged.extend_from_slice(&my_id.to_le_bytes());
tagged.extend_from_slice(env);
tagged
};
let target = {
let mut map = sessions.lock().await;
map.get_mut(&session_id).and_then(|session| {
session.last_seen = Instant::now();
session.host.as_ref().map(|h| h.tx.clone())
})
};
let Some(target) = target else { continue };
match target.try_send(framed) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
info!("Relay: disconnecting a congested viewer in session '{session_id}'");
break;
}
Err(mpsc::error::TrySendError::Closed(_)) => {}
}
}
queued_bytes += len;
if queued_bytes > MAX_PEER_QUEUE_BYTES {
warn!("Relay: peer {peer} queued {queued_bytes} bytes without draining; dropping it");
break;
}
}
outbound = rx.recv() => {
let Some(bytes) = outbound else { break };
queued_bytes = queued_bytes.saturating_sub(bytes.len());
if write_half.write_all(&bytes).await.is_err() {
break;
}
}
}
}
{
let mut map = sessions.lock().await;
if let Some(session) = map.get_mut(&session_id) {
if let Some(entry) = peer_entry.as_mut() {
match peer_role {
RelayRole::Host => {
if session.host.as_ref().is_some_and(|h| h.gen == entry.gen) {
session.host = None;
}
}
RelayRole::Viewer => {
session.viewers.retain(|v| v.gen != entry.gen);
}
}
}
let left_id = if peer_role == RelayRole::Viewer {
Some(peer_entry.as_ref().map(|e| e.id).unwrap_or(0))
} else {
None
};
if let Some(host) = session.host.as_ref() {
if let Some(id) = left_id {
let mut frame = Vec::with_capacity(9);
frame.extend_from_slice(&9u32.to_le_bytes());
frame.extend_from_slice(&RELAY_CONTROL_ID.to_le_bytes());
frame.push(K_CONTROL_VIEWER_LEFT);
frame.extend_from_slice(&id.to_le_bytes());
let _ = host.tx.try_send(frame);
}
}
if session.host.is_none() && peer_role == RelayRole::Host {
session.viewers.clear();
}
if session.host.is_none() && session.viewers.is_empty() {
map.remove(&session_id);
}
}
}
info!("Relay: {:?} left session '{session_id}'", peer_role);
Ok(())
}
async fn read_frame<R: tokio::io::AsyncRead + Unpin>(reader: &mut R) -> Result<Option<Vec<u8>>> {
let mut len_buf = [0u8; 4];
match reader.read_exact(&mut len_buf).await {
Ok(_) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
Err(e) => return Err(e.into()),
}
let len = read_len_prefix(&len_buf)?;
let mut payload = vec![0u8; len];
reader.read_exact(&mut payload).await?;
let mut framed = Vec::with_capacity(4 + len);
framed.extend_from_slice(&len_buf);
framed.extend_from_slice(&payload);
Ok(Some(framed))
}
async fn read_registration<R: tokio::io::AsyncRead + Unpin>(
reader: &mut R,
) -> Result<RelayRegister> {
let mut len_buf = [0u8; 4];
reader.read_exact(&mut len_buf).await?;
let len = u32::from_le_bytes(len_buf) as usize;
if len == 0 || len > MAX_REGISTRATION_BYTES {
anyhow::bail!("relay registration size {len} is outside 1..={MAX_REGISTRATION_BYTES}");
}
let mut buf = vec![0u8; len];
reader.read_exact(&mut buf).await?;
if let Ok(reg) = bincode::deserialize::<RelayRegister>(&buf) {
return Ok(reg);
}
#[derive(serde::Deserialize)]
struct LegacyRegister {
session: String,
role: RelayRole,
token: String,
}
let legacy: LegacyRegister = bincode::deserialize(&buf)?;
Ok(RelayRegister {
session: legacy.session,
role: legacy.role,
token: legacy.token,
protocol_version: 0,
})
}
async fn record_and_check_rate(
failures: &Arc<Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>>,
peer: std::net::SocketAddr,
) -> bool {
let mut map = failures.lock().await;
let entry = map.entry(peer).or_insert((0, Instant::now()));
if entry.1.elapsed() > AUTH_FAILURE_WINDOW {
*entry = (0, Instant::now());
}
entry.0 += 1;
entry.0 > AUTH_FAILURES_ALLOWED
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn session_codes_are_six_unambiguous_symbols() {
let code = generate_session_code();
assert_eq!(code.len(), 6);
assert!(code
.chars()
.all(|c| "ABCDEFGHJKLMNPQRSTUVWXYZ23456789".contains(c)));
}
#[tokio::test]
async fn a_hostile_registration_length_is_refused_without_allocating() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&u32::MAX.to_le_bytes());
let mut cursor: &[u8] = &bytes;
let err = read_registration(&mut cursor)
.await
.unwrap_err()
.to_string();
assert!(
err.contains("max_message_size") || err.contains("outside"),
"unhelpful: {err}"
);
}
#[tokio::test]
async fn a_hostile_frame_length_is_refused_without_allocating() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&u32::MAX.to_le_bytes());
let mut cursor: &[u8] = &bytes;
let err = read_frame(&mut cursor).await.unwrap_err().to_string();
assert!(err.contains("max_message_size"), "unhelpful: {err}");
}
#[tokio::test]
async fn host_leg_frames_roundtrip_through_the_relay_transform() {
let envelope = crate::network::Message::KeepAlive { rev: 7 }
.encode()
.unwrap();
let mut viewer_leg = Vec::new();
viewer_leg.extend_from_slice(&(envelope.len() as u32).to_le_bytes());
viewer_leg.extend_from_slice(&envelope);
let tagged = tag_host_frame(3, &envelope);
let mut cursor: &[u8] = &tagged;
let (id, env) = read_host_frame(&mut cursor)
.await
.unwrap()
.expect("a tagged frame must parse");
assert_eq!(id, 3);
assert_eq!(env, envelope);
let mut out = Vec::with_capacity(4 + env.len());
out.extend_from_slice(&(env.len() as u32).to_le_bytes());
out.extend_from_slice(&env);
assert_eq!(out, viewer_leg);
}
#[tokio::test]
async fn a_host_leg_frame_too_short_for_an_id_is_refused() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&0u32.to_le_bytes());
let mut cursor: &[u8] = &bytes;
let err = read_host_frame(&mut cursor).await.unwrap_err().to_string();
assert!(err.contains("peer id"), "unhelpful: {err}");
}
#[tokio::test]
async fn rate_limiting_trips_after_the_configured_failures() {
let failures: Arc<Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>> =
Arc::new(Mutex::new(HashMap::new()));
let peer: std::net::SocketAddr = "10.0.0.1:5000".parse().unwrap();
for i in 0..AUTH_FAILURES_ALLOWED {
assert!(
!record_and_check_rate(&failures, peer).await,
"allowed {i} of {AUTH_FAILURES_ALLOWED} attempts"
);
}
assert!(
record_and_check_rate(&failures, peer).await,
"the limit must trip"
);
}
#[tokio::test]
async fn legacy_registration_decodes_with_version_zero() {
#[derive(serde::Serialize)]
struct Legacy {
session: String,
role: RelayRole,
token: String,
}
let body = bincode::serialize(&Legacy {
session: "ABC123".into(),
role: RelayRole::Host,
token: "tok".into(),
})
.unwrap();
let mut framed = (body.len() as u32).to_le_bytes().to_vec();
framed.extend_from_slice(&body);
let mut cursor: &[u8] = &framed;
let reg = read_registration(&mut cursor).await.unwrap();
assert_eq!(reg.protocol_version, 0);
assert_eq!(reg.role, RelayRole::Host);
}
#[tokio::test]
async fn current_registration_carries_the_protocol_version() {
let reg = RelayRegister {
session: "ABC123".into(),
role: RelayRole::Viewer,
token: "tok".into(),
protocol_version: crate::network::PROTOCOL_VERSION,
};
let body = bincode::serialize(®).unwrap();
let mut framed = (body.len() as u32).to_le_bytes().to_vec();
framed.extend_from_slice(&body);
let mut cursor: &[u8] = &framed;
let back: RelayRegister = read_registration(&mut cursor).await.unwrap();
assert_eq!(back.protocol_version, crate::network::PROTOCOL_VERSION);
}
}