use super::auth;
use super::did_resolver::DidResolver;
use super::{PeerEntry, RelayState};
use crate::protocol::{SignalEnvelope, SignalPayload};
use axum::{
extract::{
OriginalUri, State, WebSocketUpgrade,
ws::{Message, WebSocket},
},
http::{HeaderMap, StatusCode},
response::IntoResponse,
};
use futures_util::{SinkExt, StreamExt};
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::{Notify, mpsc};
use uuid::Uuid;
struct AtomicGuard(Arc<AtomicUsize>);
impl Drop for AtomicGuard {
fn drop(&mut self) {
self.0.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
}
}
const RELAY_CHANNEL_CAPACITY: usize = 256;
const MAX_WS_MESSAGE_SIZE: usize = 64 * 1024;
const MAX_INVALID_MESSAGES: usize = 10;
const MAX_BACKPRESSURE_STRIKES: usize = 50;
const MAX_BACKPRESSURE_STRIKE_CAP: usize = MAX_BACKPRESSURE_STRIKES * 2;
const SILENCER_REARM_THRESHOLD: usize = 0;
const BACKPRESSURE_LOG_COOLDOWN: Duration = Duration::from_secs(30);
const WS_IDLE_TIMEOUT: Duration = Duration::from_secs(120);
const WS_PING_INTERVAL: Duration = Duration::from_secs(30);
const SIGNALS_PER_REMOTE_PEER: u32 = 16;
const RATE_BURST_FLOOR: u32 = 1024;
const RATE_BURST_CEILING: u32 = 16_384;
const RATE_REFILL_PER_SECOND: u32 = 20;
const PER_TARGET_BURST_LIMIT: u32 = 16;
const TARGET_KICK_STRIKES: u64 = RELAY_CHANNEL_CAPACITY as u64;
const PER_TARGET_WINDOW: Duration = Duration::from_secs(1);
const MAX_UNIQUE_TARGETS_FLOOR: usize = 256;
const MAX_UNIQUE_TARGETS_CEILING: usize = 4096;
fn rate_burst_for(max_peers: usize) -> u32 {
let scaled = (max_peers as u64).saturating_mul(SIGNALS_PER_REMOTE_PEER as u64);
let scaled = scaled.min(RATE_BURST_CEILING as u64) as u32;
scaled.max(RATE_BURST_FLOOR)
}
fn unique_targets_for(max_peers: usize) -> usize {
if max_peers == 0 {
return MAX_UNIQUE_TARGETS_CEILING;
}
max_peers.clamp(MAX_UNIQUE_TARGETS_FLOOR, MAX_UNIQUE_TARGETS_CEILING)
}
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(15);
const UNLIMITED_HANDSHAKE_BUDGET: usize = 256;
const WS_WRITE_TIMEOUT: Duration = Duration::from_secs(5);
fn capacity_snapshot(state: &RelayState) -> (usize, usize) {
let rooms = state.peers.len();
let peers = state.peers.iter().map(|r| r.value().len()).sum::<usize>();
(peers, rooms)
}
fn format_capacity_suffix(peers: usize, rooms: usize, max_peers: usize) -> String {
if max_peers == 0 {
format!("({peers} / unlimited peers connected in {rooms} active rooms)")
} else {
format!("({peers} / {max_peers} peers connected in {rooms} active rooms)")
}
}
fn remove_own_entry(state: &RelayState, room: &str, session_id: &str, conn_id: Uuid) -> bool {
let removed = state
.peers
.get(room)
.and_then(|inner| inner.remove_if(session_id, |_, entry| entry.conn_id == conn_id))
.is_some();
if removed {
state.peers.remove_if(room, |_, inner| inner.is_empty());
}
removed
}
fn collect_same_room_senders(
state: &RelayState,
room: &str,
exclude_session: &str,
) -> Vec<mpsc::Sender<SignalEnvelope>> {
let inner = match state.peers.get(room) {
Some(r) => Arc::clone(r.value()),
None => return Vec::new(),
};
inner
.iter()
.filter(|e| e.key() != exclude_session)
.map(|e| e.value().tx.clone())
.collect()
}
fn collect_same_room_session_ids(
state: &RelayState,
room: &str,
exclude_session: &str,
) -> Vec<String> {
let inner = match state.peers.get(room) {
Some(r) => Arc::clone(r.value()),
None => return Vec::new(),
};
inner
.iter()
.filter(|e| e.key() != exclude_session)
.map(|e| e.key().clone())
.collect()
}
pub async fn ws_handler(
ws: WebSocketUpgrade,
OriginalUri(uri): OriginalUri,
headers: HeaderMap,
State(state): State<RelayState>,
) -> impl IntoResponse {
let raw_room = uri.path().trim_matches('/');
let room = match percent_encoding::percent_decode_str(raw_room).decode_utf8() {
Ok(decoded) => decoded.into_owned(),
Err(_) => {
return (StatusCode::BAD_REQUEST, "room path is not valid UTF-8").into_response();
}
};
let room = if room.is_empty() {
"default".to_string()
} else {
room
};
let _conn_guard = if state.max_peers > 0 {
let prev = state
.active_connections
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if prev >= state.max_peers {
state
.active_connections
.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
return (StatusCode::SERVICE_UNAVAILABLE, "Relay is at capacity").into_response();
}
Some(AtomicGuard(Arc::clone(&state.active_connections)))
} else {
None
};
let handshake_limit = if state.max_peers > 0 {
(state.max_peers / 4).max(1)
} else {
UNLIMITED_HANDSHAKE_BUDGET
};
let _handshake_guard = {
let prev = state
.active_handshakes
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if prev >= handshake_limit {
state
.active_handshakes
.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
return (
StatusCode::SERVICE_UNAVAILABLE,
"Too many connections pending authentication",
)
.into_response();
}
AtomicGuard(Arc::clone(&state.active_handshakes))
};
let identity = match tokio::time::timeout(
HANDSHAKE_TIMEOUT,
extract_identity(
&headers,
state.auth_required,
state.did_resolver.as_ref(),
state.service_did.as_deref(),
),
)
.await
{
Ok(result) => result,
Err(_) => {
tracing::warn!("identity extraction timed out, dropping connection");
return (StatusCode::GATEWAY_TIMEOUT, "Authentication timed out").into_response();
}
};
drop(_handshake_guard);
let used_subprotocol = headers
.get("sec-websocket-protocol")
.and_then(|v| v.to_str().ok())
.map(|s| {
let parts: Vec<&str> = s.split(',').map(str::trim).collect();
parts.len() == 2 && parts[0] == "access_token"
})
.unwrap_or(false);
match identity {
Err(rejection) => {
rejection.into_response()
}
Ok(id) => {
if let Some(ref validated) = id {
let (peers, rooms) = capacity_snapshot(&state);
let suffix = format_capacity_suffix(peers, rooms, state.max_peers);
tracing::debug!(
did = %validated.did,
peers_connected = peers,
active_rooms = rooms,
max_peers = state.max_peers,
"JWT signature verified via DID document {suffix}"
);
}
let upgrade = ws.max_message_size(MAX_WS_MESSAGE_SIZE);
let upgrade = if used_subprotocol {
upgrade.protocols(["access_token"])
} else {
upgrade
};
upgrade
.on_upgrade(move |socket| handle_socket(socket, state, id, room, _conn_guard))
.into_response()
}
}
}
async fn extract_identity(
headers: &HeaderMap,
auth_required: bool,
resolver: Option<&DidResolver>,
service_did: Option<&str>,
) -> Result<Option<auth::ValidatedIdentity>, (StatusCode, &'static str)> {
let header_token = headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.strip_prefix("Bearer "));
let protocol_token = headers
.get("sec-websocket-protocol")
.and_then(|v| v.to_str().ok())
.and_then(|s| {
let parts: Vec<&str> = s.split(',').map(str::trim).collect();
if parts.len() == 2 && parts[0] == "access_token" {
Some(parts[1])
} else {
None
}
});
let token = header_token.or(protocol_token);
let Some(token) = token else {
if auth_required {
return Err((StatusCode::UNAUTHORIZED, "Authorization required"));
}
return Ok(None);
};
let Some(resolver) = resolver else {
if auth_required {
return Err((StatusCode::UNAUTHORIZED, "No DID resolver configured"));
}
tracing::debug!("ignoring JWT — no DID resolver available to verify signature");
return Ok(None);
};
match auth::validate_atproto_jwt(token, resolver, service_did).await {
Ok(identity) => Ok(Some(identity)),
Err(auth::AuthError::Transient(e)) => {
tracing::warn!(error = %e, "JWT validation failed: transient resolver error");
Err((
StatusCode::SERVICE_UNAVAILABLE,
"Service temporarily unavailable, please retry",
))
}
Err(auth::AuthError::InvalidToken(e)) => {
tracing::warn!(error = %e, "JWT validation failed");
Err((StatusCode::UNAUTHORIZED, "Invalid JWT"))
}
}
}
async fn handle_socket(
socket: WebSocket,
state: RelayState,
identity: Option<auth::ValidatedIdentity>,
room: String,
_conn_guard: Option<AtomicGuard>,
) {
let session_id = match identity {
Some(ref id) => id.did.clone(),
None => Uuid::new_v4().to_string(),
};
let conn_id = Uuid::new_v4();
let (mut ws_tx, mut ws_rx) = socket.split();
let (relay_tx, mut relay_rx) = mpsc::channel::<SignalEnvelope>(RELAY_CHANNEL_CAPACITY);
let self_backpressure_strikes = Arc::new(AtomicU64::new(0));
let self_shutdown = Arc::new(Notify::new());
let old_entry_info = {
use dashmap::mapref::entry::Entry;
let new_peer = PeerEntry {
tx: relay_tx,
conn_id,
backpressure_strikes: Arc::clone(&self_backpressure_strikes),
shutdown: Arc::clone(&self_shutdown),
};
let room_ref = state
.peers
.entry(room.clone())
.or_insert_with(|| Arc::new(dashmap::DashMap::new()));
let inner = room_ref.value();
match inner.entry(session_id.clone()) {
Entry::Occupied(mut occ) => {
let old = occ.insert(new_peer);
Some(old.tx)
}
Entry::Vacant(vac) => {
vac.insert(new_peer);
None
}
}
};
if let Some(_old_tx) = old_entry_info {
let leave_senders = collect_same_room_senders(&state, &room, &session_id);
for sender in leave_senders {
let envelope = SignalEnvelope {
peer_id: session_id.clone(),
signal: SignalPayload::PeerLeft(session_id.clone()),
};
let _ = sender.try_send(envelope);
}
}
let welcome = serde_json::json!({ "type": "session_id", "id": session_id });
if !matches!(
tokio::time::timeout(
WS_WRITE_TIMEOUT,
ws_tx.send(Message::Text(welcome.to_string().into())),
)
.await,
Ok(Ok(()))
) {
remove_own_entry(&state, &room, &session_id, conn_id);
return;
}
let existing_peers = collect_same_room_session_ids(&state, &room, &session_id);
tracing::debug!(
session = %session_id,
room = %room,
peer_count = existing_peers.len(),
peers = ?existing_peers,
"sending peer_list to new peer",
);
let peer_list = serde_json::json!({ "type": "peer_list", "peers": existing_peers });
if !matches!(
tokio::time::timeout(
WS_WRITE_TIMEOUT,
ws_tx.send(Message::Text(peer_list.to_string().into())),
)
.await,
Ok(Ok(()))
) {
remove_own_entry(&state, &room, &session_id, conn_id);
return;
}
let peer_senders = collect_same_room_senders(&state, &room, &session_id);
tracing::debug!(
session = %session_id,
room = %room,
notify_count = peer_senders.len(),
"broadcasting PeerJoined to room",
);
for sender in peer_senders {
let envelope = SignalEnvelope {
peer_id: session_id.clone(),
signal: SignalPayload::PeerJoined(session_id.clone()),
};
let _ = sender.try_send(envelope);
}
let session_id_write = session_id.clone();
let room_write = room.clone();
let state_write = state.clone();
let shutdown_for_write = Arc::clone(&self_shutdown);
let write_task = tokio::spawn(async move {
let mut ping_interval = tokio::time::interval(WS_PING_INTERVAL);
ping_interval.tick().await;
loop {
tokio::select! {
envelope = relay_rx.recv() => {
let Some(envelope) = envelope else { break };
let json = match serde_json::to_string(&envelope) {
Ok(j) => j,
Err(_) => continue,
};
if !matches!(
tokio::time::timeout(WS_WRITE_TIMEOUT, ws_tx.send(Message::Text(json.into()))).await,
Ok(Ok(()))
) {
tracing::warn!("disconnecting peer: write timeout or error");
break;
}
}
_ = ping_interval.tick() => {
if !matches!(
tokio::time::timeout(WS_WRITE_TIMEOUT, ws_tx.send(Message::Ping(vec![].into()))).await,
Ok(Ok(()))
) {
tracing::warn!("disconnecting peer: write timeout or error");
break;
}
}
_ = shutdown_for_write.notified() => {
tracing::warn!(
"disconnecting peer: sustained channel-full backpressure \
({TARGET_KICK_STRIKES} aggregate strikes)"
);
let _ = tokio::time::timeout(
WS_WRITE_TIMEOUT,
ws_tx.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: 1013, reason: "relay queue saturated".into(),
}))),
).await;
break;
}
}
}
});
let read_task = async {
let mut invalid_count: usize = 0;
#[derive(Default)]
struct BackpressureState {
strikes: usize,
silenced: bool,
last_warning_at: Option<tokio::time::Instant>,
}
let mut backpressure_strikes: std::collections::HashMap<String, BackpressureState> =
std::collections::HashMap::new();
let mut per_target_counts: std::collections::HashMap<String, u32> =
std::collections::HashMap::new();
let mut per_target_last_reset = tokio::time::Instant::now();
let rate_burst_capacity = rate_burst_for(state_write.max_peers);
let max_unique_targets = unique_targets_for(state_write.max_peers);
let mut rate_tokens: u32 = rate_burst_capacity;
let mut rate_last_refill = tokio::time::Instant::now();
loop {
let msg = match tokio::time::timeout(WS_IDLE_TIMEOUT, ws_rx.next()).await {
Ok(Some(Ok(msg))) => msg,
Ok(Some(Err(_))) => break, Ok(None) => break, Err(_) => {
tracing::info!(
session = %session_id,
timeout_secs = WS_IDLE_TIMEOUT.as_secs(),
"disconnecting idle peer (no messages received within timeout)"
);
break;
}
};
let text = match msg {
Message::Text(t) => t.to_string(),
Message::Close(_) => break,
Message::Ping(_) | Message::Pong(_) => continue,
_ => {
invalid_count += 1;
if invalid_count >= MAX_INVALID_MESSAGES {
tracing::warn!(
session = %session_id,
"disconnecting peer after {MAX_INVALID_MESSAGES} invalid messages"
);
break;
}
continue;
}
};
let envelope: SignalEnvelope = match serde_json::from_str(&text) {
Ok(e) => e,
Err(err) => {
invalid_count += 1;
tracing::warn!(
error = %err,
count = invalid_count,
"invalid signal envelope from peer"
);
if invalid_count >= MAX_INVALID_MESSAGES {
tracing::warn!(
session = %session_id,
"disconnecting peer after {MAX_INVALID_MESSAGES} invalid messages"
);
break;
}
continue;
}
};
if matches!(
envelope.signal,
SignalPayload::PeerJoined(_) | SignalPayload::PeerLeft(_)
) {
invalid_count += 1;
tracing::warn!(
session = %session_id,
count = invalid_count,
"dropping forged control signal from client"
);
if invalid_count >= MAX_INVALID_MESSAGES {
tracing::warn!(
session = %session_id,
"disconnecting peer after {MAX_INVALID_MESSAGES} invalid messages"
);
break;
}
continue;
}
let now = tokio::time::Instant::now();
let elapsed = now.saturating_duration_since(rate_last_refill);
if elapsed >= Duration::from_millis(50) {
let elapsed_ms = (elapsed.as_millis() as u64).min(u32::MAX as u64) as u32;
let refill = elapsed_ms.saturating_mul(RATE_REFILL_PER_SECOND) / 1000;
if refill > 0 {
let new_tokens = (rate_tokens + refill).min(rate_burst_capacity);
rate_tokens = new_tokens;
if new_tokens == rate_burst_capacity {
rate_last_refill = now;
} else {
rate_last_refill += Duration::from_millis(
refill as u64 * 1000 / RATE_REFILL_PER_SECOND as u64,
);
}
}
}
if now.saturating_duration_since(per_target_last_reset) >= PER_TARGET_WINDOW {
per_target_counts.clear();
per_target_last_reset = now;
}
if rate_tokens == 0 {
tracing::warn!(
session = %session_id,
burst = rate_burst_capacity,
refill_per_sec = RATE_REFILL_PER_SECOND,
"disconnecting peer: signal rate budget exhausted"
);
break;
}
rate_tokens -= 1;
let target_id = envelope.peer_id.clone();
if target_id == session_id {
tracing::warn!(
session = %session_id,
"dropping self-targeted signal"
);
continue;
}
if !per_target_counts.contains_key(&target_id)
&& per_target_counts.len() >= max_unique_targets
{
tracing::warn!(
session = %session_id,
unique_targets = per_target_counts.len(),
cap = max_unique_targets,
"disconnecting peer: exceeded unique target cap"
);
break;
}
let target_count = per_target_counts.entry(target_id.clone()).or_insert(0);
if *target_count >= PER_TARGET_BURST_LIMIT {
tracing::debug!(
session = %session_id,
target = %target_id,
limit = PER_TARGET_BURST_LIMIT,
"per-target burst limit reached, dropping signal"
);
continue;
}
*target_count += 1;
let signal_kind = match &envelope.signal {
SignalPayload::Offer(_) => "Offer",
SignalPayload::Answer(_) => "Answer",
SignalPayload::IceCandidate(_) => "IceCandidate",
SignalPayload::PeerJoined(_) => "PeerJoined",
SignalPayload::PeerLeft(_) => "PeerLeft",
};
tracing::debug!(
sender = %session_id,
target = %target_id,
room = %room,
signal = signal_kind,
"relay forwarding signal",
);
let forwarded = SignalEnvelope {
peer_id: session_id.clone(),
signal: envelope.signal,
};
let target_entry = state_write.peers.get(&room).and_then(|inner| {
inner.get(&target_id).map(|e| {
let entry = e.value();
(
entry.tx.clone(),
Arc::clone(&entry.backpressure_strikes),
Arc::clone(&entry.shutdown),
)
})
});
if let Some((target_tx, target_strikes, target_shutdown)) = target_entry {
match target_tx.try_send(forwarded) {
Ok(()) => {
let _ = target_strikes.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|v| Some(v.saturating_sub(1)),
);
if let Some(state) = backpressure_strikes.get_mut(&target_id) {
state.strikes = state.strikes.saturating_sub(1);
if state.strikes == SILENCER_REARM_THRESHOLD {
let cooldown_expired = state
.last_warning_at
.map(|t| {
now.saturating_duration_since(t)
>= BACKPRESSURE_LOG_COOLDOWN
})
.unwrap_or(true);
if cooldown_expired {
backpressure_strikes.remove(&target_id);
}
}
}
}
Err(mpsc::error::TrySendError::Full(_)) => {
let prev = target_strikes.fetch_add(1, Ordering::Relaxed);
if prev + 1 >= TARGET_KICK_STRIKES {
target_shutdown.notify_waiters();
}
let state = backpressure_strikes.entry(target_id.clone()).or_default();
state.strikes = state
.strikes
.saturating_add(1)
.min(MAX_BACKPRESSURE_STRIKE_CAP);
let cooldown_elapsed = state
.last_warning_at
.map(|t| now.saturating_duration_since(t) >= BACKPRESSURE_LOG_COOLDOWN)
.unwrap_or(true);
if cooldown_elapsed && !state.silenced {
tracing::warn!(
target = %target_id,
sender = %session_id,
"relay channel full for target, dropping signal (backpressure)"
);
state.last_warning_at = Some(now);
}
if !state.silenced && state.strikes >= MAX_BACKPRESSURE_STRIKES {
state.silenced = true;
tracing::warn!(
target = %target_id,
sender = %session_id,
"target hit {MAX_BACKPRESSURE_STRIKES} backpressure \
strikes; silencing logs until channel drains"
);
}
}
Err(mpsc::error::TrySendError::Closed(_)) => {
backpressure_strikes.remove(&target_id);
tracing::debug!(
target = %target_id,
"target channel closed, dropping this message"
);
}
}
} else {
backpressure_strikes.remove(&target_id);
tracing::debug!(target = %target_id, "target peer not found, dropping signal");
}
}
};
let mut write_task = write_task;
tokio::select! {
_ = &mut write_task => {
tracing::debug!(session = %session_id_write, "write task ended, stopping read");
}
_ = read_task => {
tracing::debug!(session = %session_id_write, "read task ended, aborting write task");
write_task.abort();
}
}
let was_removed = remove_own_entry(&state, &room_write, &session_id_write, conn_id);
if !was_removed {
tracing::debug!(
session = %session_id_write,
room = %room_write,
"stale connection cleanup skipped (session was replaced)"
);
return;
}
let remaining_senders = collect_same_room_senders(&state, &room_write, &session_id_write);
for sender in remaining_senders {
let envelope = SignalEnvelope {
peer_id: session_id_write.clone(),
signal: SignalPayload::PeerLeft(session_id_write.clone()),
};
let _ = sender.try_send(envelope);
}
drop(_conn_guard);
let (peers, rooms) = capacity_snapshot(&state);
let suffix = format_capacity_suffix(peers, rooms, state.max_peers);
tracing::info!(
session = %session_id_write,
peers_connected = peers,
active_rooms = rooms,
max_peers = state.max_peers,
"peer disconnected {suffix}"
);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rate_burst_floor_applies_for_small_relays() {
assert_eq!(rate_burst_for(1), RATE_BURST_FLOOR);
assert_eq!(rate_burst_for(8), RATE_BURST_FLOOR);
}
#[test]
fn rate_burst_scales_with_max_peers() {
let burst = rate_burst_for(512);
assert!(burst >= 512 * SIGNALS_PER_REMOTE_PEER);
assert!(burst <= RATE_BURST_CEILING);
assert!(burst > 500);
}
#[test]
fn rate_burst_caps_at_ceiling() {
assert_eq!(rate_burst_for(1_000_000), RATE_BURST_CEILING);
}
#[test]
fn rate_burst_for_unlimited_uses_floor() {
assert_eq!(rate_burst_for(0), RATE_BURST_FLOOR);
}
#[test]
fn unique_targets_cover_default_max_peers() {
assert!(unique_targets_for(512) >= 512);
}
#[test]
fn unique_targets_floor_for_small_rooms() {
assert_eq!(unique_targets_for(8), MAX_UNIQUE_TARGETS_FLOOR);
assert_eq!(unique_targets_for(256), MAX_UNIQUE_TARGETS_FLOOR);
}
#[test]
fn unique_targets_caps_at_ceiling() {
assert_eq!(unique_targets_for(1_000_000), MAX_UNIQUE_TARGETS_CEILING);
}
#[test]
fn unique_targets_unlimited_uses_ceiling() {
assert_eq!(unique_targets_for(0), MAX_UNIQUE_TARGETS_CEILING);
}
#[test]
fn per_target_burst_limit_well_below_channel_capacity() {
assert!(
RELAY_CHANNEL_CAPACITY as u32 / PER_TARGET_BURST_LIMIT >= 16,
"PER_TARGET_BURST_LIMIT must leave at least 16× headroom below \
RELAY_CHANNEL_CAPACITY so a single sender cannot fill the channel"
);
}
#[test]
fn target_kick_threshold_bounds_single_sender_contribution() {
assert!(
TARGET_KICK_STRIKES > PER_TARGET_BURST_LIMIT as u64,
"TARGET_KICK_STRIKES ({TARGET_KICK_STRIKES}) must exceed one sender's \
contribution ({PER_TARGET_BURST_LIMIT}) so no single peer can kick a target"
);
}
}