use super::auth;
use super::did_resolver::DidResolver;
use super::{PeerEntry, RelayState};
use crate::protocol::{SignalEnvelope, SignalPayload};
use axum::{
extract::{
OriginalUri, Query, State, WebSocketUpgrade,
ws::{Message, WebSocket},
},
http::{HeaderMap, StatusCode},
response::IntoResponse,
};
use futures_util::{SinkExt, StreamExt};
use serde::Deserialize;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::time::Duration;
use tokio::sync::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 WS_IDLE_TIMEOUT: Duration = Duration::from_secs(120);
const WS_PING_INTERVAL: Duration = Duration::from_secs(30);
const RATE_BURST_CAPACITY: u32 = 500;
const RATE_REFILL_PER_SECOND: u32 = 20;
const PER_TARGET_BURST_LIMIT: u32 = 64;
const PER_TARGET_WINDOW: Duration = Duration::from_secs(1);
const MAX_UNIQUE_TARGETS: usize = 256;
const MAX_PEER_ID_LENGTH: usize = 512;
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(15);
const WS_WRITE_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Debug, Deserialize)]
pub struct WsQueryParams {
token: Option<String>,
}
pub async fn ws_handler(
ws: WebSocketUpgrade,
OriginalUri(uri): OriginalUri,
headers: HeaderMap,
Query(query): Query<WsQueryParams>,
State(state): State<RelayState>,
) -> impl IntoResponse {
let room = uri.path().trim_matches('/').to_string();
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_guard = if state.max_peers > 0 {
let handshake_limit = (state.max_peers / 4).max(1);
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();
}
Some(AtomicGuard(Arc::clone(&state.active_handshakes)))
} else {
None
};
let identity = match tokio::time::timeout(
HANDSHAKE_TIMEOUT,
extract_identity(
&headers,
query.token.as_deref(),
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) => {
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,
query_token: Option<&str>,
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).or(query_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 old_entry_info = {
use dashmap::mapref::entry::Entry;
let new_peer = PeerEntry {
tx: relay_tx,
conn_id,
room: room.clone(),
};
match state.peers.entry(session_id.clone()) {
Entry::Occupied(mut occ) => {
let old = occ.insert(new_peer);
Some((old.room, old.tx))
}
Entry::Vacant(vac) => {
vac.insert(new_peer);
None
}
}
};
if let Some((old_room, _old_tx)) = old_entry_info {
let leave_senders: Vec<mpsc::Sender<SignalEnvelope>> = state
.peers
.iter()
.filter(|entry| *entry.key() != session_id && entry.value().room == old_room)
.map(|entry| entry.value().tx.clone())
.collect();
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(()))
) {
state
.peers
.remove_if(&session_id, |_, entry| entry.conn_id == conn_id);
return;
}
let existing_peers: Vec<String> = state
.peers
.iter()
.filter(|entry| *entry.key() != session_id && entry.value().room == room)
.map(|entry| entry.key().clone())
.collect();
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(()))
) {
state
.peers
.remove_if(&session_id, |_, entry| entry.conn_id == conn_id);
return;
}
let peer_senders: Vec<mpsc::Sender<SignalEnvelope>> = state
.peers
.iter()
.filter(|entry| *entry.key() != session_id && entry.value().room == room)
.map(|entry| entry.value().tx.clone())
.collect();
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 state_write = state.clone();
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;
}
}
}
}
});
let read_task = async {
let mut invalid_count: usize = 0;
let mut backpressure_strikes: std::collections::HashMap<String, usize> =
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 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 envelope.peer_id.len() > MAX_PEER_ID_LENGTH {
invalid_count += 1;
tracing::warn!(
session = %session_id,
peer_id_len = envelope.peer_id.len(),
count = invalid_count,
"dropping signal with oversized peer_id (max {MAX_PEER_ID_LENGTH} bytes)"
);
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.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.duration_since(per_target_last_reset) >= PER_TARGET_WINDOW {
per_target_counts.clear();
per_target_last_reset = now;
}
if rate_tokens == 0 {
tracing::debug!(
session = %session_id,
"dropping signal: rate limit exhausted (burst {RATE_BURST_CAPACITY}, \
refill {RATE_REFILL_PER_SECOND}/s)"
);
continue;
}
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(),
"disconnecting peer: exceeded unique target cap ({MAX_UNIQUE_TARGETS})"
);
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 forwarded = SignalEnvelope {
peer_id: session_id.clone(),
signal: envelope.signal,
};
let at_strike_limit = backpressure_strikes
.get(&target_id)
.is_some_and(|s| *s >= MAX_BACKPRESSURE_STRIKES);
if let Some(entry) = state_write.peers.get(&target_id) {
if entry.value().room != room {
tracing::warn!(
session = %session_id,
target = %target_id,
"dropping cross-room signal"
);
} else {
match entry.value().tx.try_send(forwarded) {
Ok(()) => {
backpressure_strikes.remove(&target_id);
}
Err(mpsc::error::TrySendError::Full(_)) => {
if !at_strike_limit {
let strikes =
backpressure_strikes.entry(target_id.clone()).or_insert(0);
*strikes += 1;
tracing::warn!(
target = %target_id,
sender = %session_id,
strikes = *strikes,
"relay channel full for target, dropping signal (backpressure)"
);
if *strikes >= MAX_BACKPRESSURE_STRIKES {
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 = state
.peers
.remove_if(&session_id_write, |_, entry| entry.conn_id == conn_id);
if was_removed.is_none() {
tracing::debug!(
session = %session_id_write,
"stale connection cleanup skipped (session was replaced)"
);
return;
}
let remaining_senders: Vec<mpsc::Sender<SignalEnvelope>> = state
.peers
.iter()
.filter(|entry| entry.value().room == room)
.map(|entry| entry.value().tx.clone())
.collect();
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);
}
tracing::info!(session = %session_id_write, "peer disconnected");
}