use std::borrow::Cow;
use std::collections::{HashMap, HashSet, VecDeque};
use std::fmt;
use std::net::IpAddr;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, RwLock, RwLockReadGuard, RwLockWriteGuard, Weak};
use std::time::{Duration, Instant};
use axum::extract::ws::Utf8Bytes;
use http::StatusCode;
use net_backend_protocol::{codes, CloseCode, ServerPush, UnixMillis, UserId, WsPushFrame};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use tokio::sync::{mpsc, watch};
use tokio::task::JoinHandle;
use super::handlers::HandlerMap;
use crate::auth::events::Revocation;
use crate::auth::{AuthService, Authenticator};
use crate::config::WsConfig;
use crate::error::AppError;
use crate::rate_limit::{ip_key, KeyedBuckets, RateDecision, DEFAULT_IPV6_PREFIX};
use crate::state::AppState;
pub const MAX_ROOM_NAME_BYTES: usize = 128;
const RECENT_REVOCATIONS: Duration = Duration::from_secs(30);
const LAG_OVERLAP_MS: i64 = 10_000;
const ROLE_REFRESH_CHUNK: usize = 500;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
pub struct ConnectionId(u64);
impl ConnectionId {
pub const fn get(self) -> u64 {
self.0
}
#[doc(hidden)]
pub const fn for_tests(n: u64) -> Self {
Self(n)
}
}
impl fmt::Display for ConnectionId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&self.0, f)
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[non_exhaustive]
pub enum Target {
Connection(ConnectionId),
User(UserId),
Room(Arc<str>),
All,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Control {
LeaveRoom(Arc<str>),
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct Delivery {
pub target: Target,
pub frame: Utf8Bytes,
pub control: Option<Control>,
}
impl Delivery {
pub fn new(target: Target, frame: impl Into<Utf8Bytes>) -> Self {
Self { target, frame: frame.into(), control: None }
}
pub fn control(target: Target, control: Control) -> Self {
Self { target, frame: Utf8Bytes::from_static(""), control: Some(control) }
}
}
#[derive(Serialize)]
struct DeliveryOut<'a> {
target: &'a Target,
frame: &'a str,
#[serde(skip_serializing_if = "Option::is_none")]
control: Option<&'a Control>,
}
#[derive(Deserialize)]
struct DeliveryIn {
target: Target,
frame: String,
#[serde(default)]
control: Option<Control>,
}
impl Serialize for Delivery {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
DeliveryOut { target: &self.target, frame: self.frame.as_str(), control: self.control.as_ref() }.serialize(serializer)
}
}
impl<'de> Deserialize<'de> for Delivery {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let raw = DeliveryIn::deserialize(deserializer)?;
let mut delivery = Delivery::new(raw.target, raw.frame);
delivery.control = raw.control;
Ok(delivery)
}
}
pub trait Broadcaster: Send + Sync + 'static {
fn start(&self, local: LocalDelivery) {
let _ = local;
}
fn publish(&self, local: &LocalDelivery, delivery: Delivery);
}
#[derive(Clone, Copy, Debug, Default)]
pub struct LocalBroadcaster;
impl Broadcaster for LocalBroadcaster {
fn publish(&self, local: &LocalDelivery, delivery: Delivery) {
local.deliver(&delivery);
}
}
#[derive(Clone)]
pub struct LocalDelivery(Weak<Inner>);
impl fmt::Debug for LocalDelivery {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("LocalDelivery")
}
}
impl LocalDelivery {
pub fn deliver(&self, delivery: &Delivery) -> usize {
match self.0.upgrade() {
Some(inner) => Hub(inner).deliver_local(delivery),
None => 0,
}
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum PushError {
ReservedKind,
Encode(serde_json::Error),
#[non_exhaustive]
TooLarge {
size: usize,
limit: usize,
},
}
impl fmt::Display for PushError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
PushError::ReservedKind => f.write_str("a push kind must not be empty or an auth kind"),
PushError::Encode(error) => write!(f, "the push could not be encoded: {error}"),
PushError::TooLarge { size, limit } => write!(f, "the push is {size} bytes, over ws.max_message_bytes ({limit})"),
}
}
}
impl std::error::Error for PushError {}
impl From<PushError> for AppError {
fn from(error: PushError) -> Self {
AppError::internal(error)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum JoinError {
NotConnected,
RoomFull,
TooManyRooms,
InvalidRoom,
}
impl fmt::Display for JoinError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
JoinError::NotConnected => "the connection is gone",
JoinError::RoomFull => "the room is full",
JoinError::TooManyRooms => "too many rooms joined on this connection",
JoinError::InvalidRoom => "invalid room name",
})
}
}
impl std::error::Error for JoinError {}
impl From<JoinError> for AppError {
fn from(error: JoinError) -> Self {
match error {
JoinError::RoomFull => AppError::new(codes::ROOM_FULL, "the room is full"),
JoinError::TooManyRooms => AppError::new(codes::QUOTA_EXCEEDED, "too many rooms joined on this connection"),
JoinError::InvalidRoom => AppError::bad_request("invalid room name"),
JoinError::NotConnected => AppError::with_status(StatusCode::SERVICE_UNAVAILABLE, codes::UNAVAILABLE, "the connection is gone"),
}
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct ConnectionInfo {
pub id: ConnectionId,
pub user_id: Option<UserId>,
pub session_id: Option<i64>,
pub roles: Vec<String>,
pub ip: Option<IpAddr>,
pub connected_at: UnixMillis,
pub rooms: Vec<Arc<str>>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct HubStats {
pub connections: usize,
pub pending: usize,
pub authenticated: usize,
pub users: usize,
pub rooms: usize,
}
tokio::task_local! {
pub(crate) static HANDLING: ConnectionId;
}
pub(crate) struct Outgoing {
pub(crate) frame: Utf8Bytes,
pub(crate) after_answer: bool,
}
pub(crate) type CloseRequest = Option<(CloseCode, Cow<'static, str>)>;
pub(crate) fn sendable(code: CloseCode) -> CloseCode {
match code.get() {
1000..=1003 | 1007..=1014 | 3000..=4999 => code,
_ => CloseCode::INTERNAL_ERROR,
}
}
struct Conn {
outbox: mpsc::Sender<Outgoing>,
close: watch::Sender<CloseRequest>,
user: Option<UserId>,
session: Option<i64>,
roles: Vec<String>,
ip: Option<IpAddr>,
rooms: Vec<Arc<str>>,
connected_at: UnixMillis,
}
impl Conn {
fn request_close(&self, code: CloseCode, reason: Cow<'static, str>) -> bool {
let code = sendable(code);
self.close.send_if_modified(|current| {
if current.is_none() {
*current = Some((code, reason));
true
} else {
false
}
})
}
}
#[derive(Default)]
struct Registry {
conns: HashMap<ConnectionId, Conn>,
users: HashMap<UserId, Vec<ConnectionId>>,
rooms: HashMap<Arc<str>, HashSet<ConnectionId>>,
recent: VecDeque<(Instant, Revocation)>,
}
pub(crate) type RoomLeft = Arc<dyn Fn(&Hub, ConnectionId, UserId, &str) + Send + Sync>;
pub(crate) struct Inner {
config: WsConfig,
room_left: RwLock<Vec<RoomLeft>>,
registry: RwLock<Registry>,
next_id: AtomicU64,
live: AtomicUsize,
pending: AtomicUsize,
per_ip: Mutex<HashMap<IpAddr, usize>>,
closing: AtomicBool,
deadline: Mutex<Option<Instant>>,
kill: watch::Sender<bool>,
broadcaster: Arc<dyn Broadcaster>,
pub(crate) handlers: HandlerMap,
pub(crate) authenticators: Arc<[Arc<dyn Authenticator>]>,
handshakes: Option<KeyedBuckets<IpAddr>>,
metrics: bool,
}
#[derive(Clone)]
pub struct Hub(pub(crate) Arc<Inner>);
impl fmt::Debug for Hub {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Hub").field("stats", &self.stats()).field("kinds", &self.0.handlers.len()).finish_non_exhaustive()
}
}
pub(crate) struct Slot {
hub: Hub,
ip: Option<IpAddr>,
pending: AtomicBool,
}
impl Slot {
pub(crate) fn authenticated(&self) {
if self.pending.swap(false, Ordering::SeqCst) {
self.hub.0.pending.fetch_sub(1, Ordering::SeqCst);
}
}
}
impl Drop for Slot {
fn drop(&mut self) {
self.authenticated();
if let Some(ip) = self.ip {
let mut per_ip = self.hub.0.per_ip.lock().unwrap_or_else(|p| p.into_inner());
if let Some(n) = per_ip.get_mut(&ip) {
*n = n.saturating_sub(1);
if *n == 0 {
per_ip.remove(&ip);
}
}
}
let left = self.hub.0.live.fetch_sub(1, Ordering::SeqCst).saturating_sub(1);
if self.hub.0.metrics {
metrics::gauge!("nbs_ws_connections").set(left as f64);
}
}
}
pub(crate) enum Refusal {
Full,
TooManyPending,
TooManyFromAddress,
}
pub(crate) struct Registered {
pub(crate) id: ConnectionId,
hub: Hub,
pub(crate) slot: Slot,
}
impl Drop for Registered {
fn drop(&mut self) {
self.hub.unregister(self.id);
}
}
pub(crate) struct ConnHandle {
pub(crate) outbox: mpsc::Receiver<Outgoing>,
pub(crate) close: watch::Receiver<CloseRequest>,
pub(crate) kill: watch::Receiver<bool>,
}
fn valid_room(room: &str) -> bool {
!room.is_empty() && room.len() <= MAX_ROOM_NAME_BYTES && !room.chars().any(char::is_control)
}
impl Hub {
pub(crate) fn new(
config: WsConfig,
handlers: HandlerMap,
authenticators: Arc<[Arc<dyn Authenticator>]>,
broadcaster: Arc<dyn Broadcaster>,
metrics: bool,
) -> Self {
let handshakes =
(config.handshakes_per_ip_per_minute > 0).then(|| KeyedBuckets::new(config.handshakes_per_ip_per_minute, Duration::from_secs(60), 100_000));
Hub(Arc::new(Inner {
config,
room_left: RwLock::new(Vec::new()),
registry: RwLock::new(Registry::default()),
next_id: AtomicU64::new(1),
live: AtomicUsize::new(0),
pending: AtomicUsize::new(0),
per_ip: Mutex::new(HashMap::new()),
closing: AtomicBool::new(false),
deadline: Mutex::new(None),
kill: watch::channel(false).0,
broadcaster,
handlers,
authenticators,
handshakes,
metrics,
}))
}
fn read(&self) -> RwLockReadGuard<'_, Registry> {
self.0.registry.read().unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn write(&self) -> RwLockWriteGuard<'_, Registry> {
self.0.registry.write().unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub fn config(&self) -> &WsConfig {
&self.0.config
}
pub fn local(&self) -> LocalDelivery {
LocalDelivery(Arc::downgrade(&self.0))
}
pub(crate) fn metrics(&self) -> bool {
self.0.metrics
}
pub fn encode_push<P: ServerPush>(push: &P) -> Result<Utf8Bytes, PushError> {
Self::encode_raw(P::KIND, push)
}
pub fn encode_raw<T: Serialize + ?Sized>(kind: &str, data: &T) -> Result<Utf8Bytes, PushError> {
let frame = WsPushFrame::checked(kind, data).ok_or(PushError::ReservedKind)?;
serde_json::to_string(&frame).map(Utf8Bytes::from).map_err(PushError::Encode)
}
fn check_size(&self, frame: &Utf8Bytes) -> Result<(), PushError> {
let (size, limit) = (frame.as_str().len(), self.0.config.max_message_bytes);
if size > limit {
tracing::warn!(size, limit, "WebSocket push refused: larger than ws.max_message_bytes");
if self.0.metrics {
metrics::counter!("nbs_ws_dropped_frames_total", "reason" => "push_too_big").increment(1);
}
return Err(PushError::TooLarge { size, limit });
}
Ok(())
}
pub fn push<P: ServerPush>(&self, target: Target, push: &P) -> Result<(), PushError> {
self.publish(Delivery::new(target, Self::encode_push(push)?))
}
pub fn push_user<P: ServerPush>(&self, user: UserId, push: &P) -> Result<(), PushError> {
self.push(Target::User(user), push)
}
pub fn push_room<P: ServerPush>(&self, room: &str, push: &P) -> Result<(), PushError> {
self.push(Target::Room(Arc::from(room)), push)
}
pub fn push_all<P: ServerPush>(&self, push: &P) -> Result<(), PushError> {
self.push(Target::All, push)
}
pub fn push_connection<P: ServerPush>(&self, connection: ConnectionId, push: &P) -> Result<(), PushError> {
self.push(Target::Connection(connection), push)
}
pub fn push_raw<T: Serialize + ?Sized>(&self, target: Target, kind: &str, data: &T) -> Result<(), PushError> {
self.publish(Delivery::new(target, Self::encode_raw(kind, data)?))
}
pub fn publish(&self, delivery: Delivery) -> Result<(), PushError> {
self.check_size(&delivery.frame)?;
if matches!(delivery.target, Target::Connection(_)) {
self.deliver_local(&delivery);
} else {
self.0.broadcaster.publish(&self.local(), delivery);
}
Ok(())
}
fn deliver_local(&self, delivery: &Delivery) -> usize {
if let Some(control) = &delivery.control {
return self.apply_control(&delivery.target, control);
}
if self.check_size(&delivery.frame).is_err() {
return 0;
}
let handling = HANDLING.try_with(|id| *id).ok();
let registry = self.read();
let mut sent = 0usize;
let mut slow = 0usize;
let mut send = |(id, conn): (&ConnectionId, &Conn)| {
if conn.user.is_none() {
return;
}
let outgoing = Outgoing { frame: delivery.frame.clone(), after_answer: handling == Some(*id) };
match conn.outbox.try_send(outgoing) {
Ok(()) => sent += 1,
Err(mpsc::error::TrySendError::Full(_)) => {
if conn.request_close(CloseCode::TRY_AGAIN_LATER, Cow::Borrowed("too slow to read its messages")) {
slow += 1;
}
}
Err(mpsc::error::TrySendError::Closed(_)) => {}
}
};
match &delivery.target {
Target::Connection(id) => registry.conns.get_key_value(id).into_iter().for_each(&mut send),
Target::User(user) => registry.users.get(user).into_iter().flatten().filter_map(|id| registry.conns.get_key_value(id)).for_each(&mut send),
Target::Room(room) => registry.rooms.get(room).into_iter().flatten().filter_map(|id| registry.conns.get_key_value(id)).for_each(&mut send),
Target::All => registry.conns.iter().for_each(&mut send),
}
drop(registry);
if slow > 0 {
tracing::warn!(sockets = slow, "WebSocket outbox full: closing slow sockets with 1013");
if self.0.metrics {
metrics::counter!("nbs_ws_slow_consumers_total").increment(slow as u64);
}
}
sent
}
fn apply_control(&self, target: &Target, control: &Control) -> usize {
match control {
Control::LeaveRoom(room) => {
let ids: Vec<ConnectionId> = {
let registry = self.read();
match target {
Target::Connection(id) => vec![*id],
Target::User(user) => registry.users.get(user).cloned().unwrap_or_default(),
Target::Room(name) => registry.rooms.get(name).map(|m| m.iter().copied().collect()).unwrap_or_default(),
Target::All => registry.conns.keys().copied().collect(),
}
};
let mut left = Vec::new();
for id in ids {
if self.leave(id, room) {
if let Some(user) = self.read().conns.get(&id).and_then(|c| c.user) {
left.push((id, user));
}
}
}
let listeners: Vec<RoomLeft> = self.0.room_left.read().unwrap_or_else(|p| p.into_inner()).clone();
for (id, user) in &left {
for listener in &listeners {
listener(self, *id, *user, room);
}
}
left.len()
}
}
}
#[cfg_attr(not(feature = "chat"), allow(dead_code))]
pub(crate) fn on_room_left(&self, listener: RoomLeft) {
self.0.room_left.write().unwrap_or_else(|p| p.into_inner()).push(listener);
}
pub fn remove_from_room(&self, user: UserId, room: &str) -> Result<(), PushError> {
self.publish(Delivery::control(Target::User(user), Control::LeaveRoom(Arc::from(room))))
}
pub fn join(&self, connection: ConnectionId, room: impl Into<Arc<str>>) -> Result<bool, JoinError> {
self.join_with_cap(connection, room, self.0.config.max_room_members)
}
pub fn join_with_cap(&self, connection: ConnectionId, room: impl Into<Arc<str>>, cap: usize) -> Result<bool, JoinError> {
let room = room.into();
if !valid_room(&room) {
return Err(JoinError::InvalidRoom);
}
let max_rooms = self.0.config.max_rooms_per_connection;
let mut registry = self.write();
let Registry { conns, rooms, .. } = &mut *registry;
let conn = conns.get_mut(&connection).filter(|c| c.user.is_some()).ok_or(JoinError::NotConnected)?;
if conn.rooms.iter().any(|r| **r == *room) {
return Ok(false);
}
if conn.rooms.len() >= max_rooms {
return Err(JoinError::TooManyRooms);
}
let members = rooms.entry(room.clone()).or_default();
if members.len() >= cap {
if members.is_empty() {
rooms.remove(&room);
}
return Err(JoinError::RoomFull);
}
members.insert(connection);
conn.rooms.push(room);
Ok(true)
}
pub fn leave(&self, connection: ConnectionId, room: &str) -> bool {
let mut registry = self.write();
let Registry { conns, rooms, .. } = &mut *registry;
let Some(conn) = conns.get_mut(&connection) else { return false };
let Some(index) = conn.rooms.iter().position(|r| &**r == room) else { return false };
conn.rooms.swap_remove(index);
if let Some(members) = rooms.get_mut(room) {
members.remove(&connection);
if members.is_empty() {
rooms.remove(room);
}
}
true
}
pub fn room_members(&self, room: &str) -> Vec<ConnectionId> {
self.read().rooms.get(room).map(|m| m.iter().copied().collect()).unwrap_or_default()
}
pub fn room_size(&self, room: &str) -> usize {
self.read().rooms.get(room).map_or(0, HashSet::len)
}
pub fn rooms_of(&self, connection: ConnectionId) -> Vec<Arc<str>> {
self.read().conns.get(&connection).map(|c| c.rooms.clone()).unwrap_or_default()
}
pub fn connection(&self, connection: ConnectionId) -> Option<ConnectionInfo> {
self.read().conns.get(&connection).map(|c| ConnectionInfo {
id: connection,
user_id: c.user,
session_id: c.session,
roles: c.roles.clone(),
ip: c.ip,
connected_at: c.connected_at,
rooms: c.rooms.clone(),
})
}
pub fn connections_of(&self, user: UserId) -> Vec<ConnectionId> {
self.read().users.get(&user).cloned().unwrap_or_default()
}
pub fn is_online(&self, user: UserId) -> bool {
self.read().users.get(&user).is_some_and(|c| !c.is_empty())
}
pub fn stats(&self) -> HubStats {
let connections = self.0.live.load(Ordering::SeqCst);
let pending = self.0.pending.load(Ordering::SeqCst);
let registry = self.read();
HubStats {
connections,
pending,
authenticated: registry.conns.values().filter(|c| c.user.is_some()).count(),
users: registry.users.len(),
rooms: registry.rooms.len(),
}
}
pub fn close(&self, connection: ConnectionId, code: CloseCode, reason: impl Into<Cow<'static, str>>) -> bool {
self.read().conns.get(&connection).is_some_and(|c| c.request_close(code, reason.into()))
}
pub fn close_user(&self, user: UserId, code: CloseCode, reason: impl Into<Cow<'static, str>>) -> usize {
let reason = reason.into();
let registry = self.read();
registry.users.get(&user).into_iter().flatten().filter_map(|id| registry.conns.get(id)).filter(|c| c.request_close(code, reason.clone())).count()
}
pub(crate) fn roles_of(&self, connection: ConnectionId) -> Option<Vec<String>> {
self.read().conns.get(&connection).map(|c| c.roles.clone())
}
pub(crate) fn set_roles(&self, user: UserId, roles: &[String]) {
let mut registry = self.write();
let Registry { conns, users, .. } = &mut *registry;
for id in users.get(&user).into_iter().flatten() {
if let Some(conn) = conns.get_mut(id) {
if conn.roles != roles {
conn.roles = roles.to_vec();
}
}
}
}
pub(crate) fn is_closing(&self) -> bool {
self.0.closing.load(Ordering::SeqCst)
}
pub(crate) fn handshake_allowed(&self, ip: Option<IpAddr>) -> RateDecision {
match (&self.0.handshakes, ip) {
(Some(buckets), Some(ip)) => buckets.check(ip_key(ip, DEFAULT_IPV6_PREFIX)),
_ => RateDecision::Allow,
}
}
pub(crate) fn try_reserve(&self, ip: Option<IpAddr>, pending: bool) -> Result<Slot, Refusal> {
let config = &self.0.config;
let key = ip.map(|ip| ip_key(ip, DEFAULT_IPV6_PREFIX));
if pending && self.0.pending.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| (n < config.max_pending_connections).then_some(n + 1)).is_err() {
return Err(Refusal::TooManyPending);
}
let undo_pending = || {
if pending {
self.0.pending.fetch_sub(1, Ordering::SeqCst);
}
};
if let Some(key) = key {
let mut per_ip = self.0.per_ip.lock().unwrap_or_else(|p| p.into_inner());
let n = per_ip.entry(key).or_insert(0);
if *n >= config.max_connections_per_ip {
if *n == 0 {
per_ip.remove(&key);
}
drop(per_ip);
undo_pending();
return Err(Refusal::TooManyFromAddress);
}
*n += 1;
}
let slot = Slot { hub: self.clone(), ip: key, pending: AtomicBool::new(pending) };
match self.0.live.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| (n < config.max_connections).then_some(n + 1)) {
Ok(reserved) => {
if self.0.metrics {
metrics::gauge!("nbs_ws_connections").set((reserved + 1) as f64);
}
Ok(slot)
}
Err(_) => {
self.0.live.fetch_add(1, Ordering::SeqCst);
drop(slot);
Err(Refusal::Full)
}
}
}
pub(crate) fn register(&self, slot: Slot, ip: Option<IpAddr>, now: UnixMillis) -> (Registered, ConnHandle) {
let id = ConnectionId(self.0.next_id.fetch_add(1, Ordering::Relaxed));
let (outbox_tx, outbox) = mpsc::channel(self.0.config.outbox_frames.max(1));
let (close_tx, close) = watch::channel(None);
let kill = self.0.kill.subscribe();
let conn = Conn { outbox: outbox_tx, close: close_tx, user: None, session: None, roles: Vec::new(), ip, rooms: Vec::new(), connected_at: now };
self.write().conns.insert(id, conn);
if self.is_closing() {
self.close(id, CloseCode::GOING_AWAY, "the server is shutting down");
}
(Registered { id, hub: self.clone(), slot }, ConnHandle { outbox, close, kill })
}
pub(crate) fn set_user(&self, id: ConnectionId, user: UserId, session: Option<i64>, roles: Vec<String>, checked_at: Instant) -> Result<(), CloseCode> {
let cap = self.0.config.max_connections_per_user.max(1);
let mut registry = self.write();
if let Some((_, revocation)) = registry.recent.iter().rev().find(|(at, r)| *at >= checked_at && r.applies_to(user, session)) {
return Err(revocation.close_code());
}
let Registry { conns, users, .. } = &mut *registry;
let Some(conn) = conns.get_mut(&id) else { return Ok(()) };
conn.user = Some(user);
conn.session = session;
conn.roles = roles;
let list = users.entry(user).or_default();
if !list.contains(&id) {
list.push(id);
}
let excess = list.len().saturating_sub(cap);
if excess > 0 {
let same = |c: &ConnectionId| conns.get(c).is_some_and(|conn| session.is_some() && conn.session == session);
let mut order: Vec<ConnectionId> = list.iter().copied().filter(|c| *c != id && same(c)).collect();
order.extend(list.iter().copied().filter(|c| *c != id && !same(c)));
for old in order.into_iter().take(excess) {
if let Some(conn) = conns.get(&old) {
conn.request_close(CloseCode::REPLACED, Cow::Borrowed("replaced by a newer connection"));
}
}
}
Ok(())
}
fn unregister(&self, id: ConnectionId) {
let mut registry = self.write();
let Some(conn) = registry.conns.remove(&id) else { return };
if let Some(user) = conn.user {
if let Some(list) = registry.users.get_mut(&user) {
list.retain(|c| *c != id);
if list.is_empty() {
registry.users.remove(&user);
}
}
}
for room in &conn.rooms {
if let Some(members) = registry.rooms.get_mut(room) {
members.remove(&id);
if members.is_empty() {
registry.rooms.remove(room);
}
}
}
}
pub(crate) fn apply_revocation(&self, revocation: &Revocation) -> usize {
let code = revocation.close_code();
let reason = if code == CloseCode::BANNED { "the account is banned" } else { "the session was revoked" };
let mut registry = self.write();
let now = Instant::now();
while registry.recent.front().is_some_and(|(at, _)| now.duration_since(*at) > RECENT_REVOCATIONS) || registry.recent.len() >= 4096 {
registry.recent.pop_front();
}
registry.recent.push_back((now, *revocation));
let closed = registry
.users
.get(&revocation.user_id)
.into_iter()
.flatten()
.filter_map(|id| registry.conns.get(id))
.filter(|c| c.user.is_some_and(|u| revocation.applies_to(u, c.session)))
.filter(|c| c.request_close(code, Cow::Borrowed(reason)))
.count();
drop(registry);
if closed > 0 {
tracing::info!(user = revocation.user_id.get(), sockets = closed, code = code.get(), "WebSocket: closing revoked sockets");
}
closed
}
async fn refresh_roles(&self, state: &AppState, service: &AuthService, users: Vec<UserId>) {
for chunk in users.chunks(ROLE_REFRESH_CHUNK) {
match service.roles_of_users(state, chunk).await {
Ok(roles) => {
for user in chunk {
self.set_roles(*user, roles.get(user).map(Vec::as_slice).unwrap_or_default());
}
}
Err(error) => tracing::warn!(%error, "WebSocket: refreshing roles failed"),
}
}
}
pub(crate) fn start(&self, state: &AppState) -> Vec<JoinHandle<()>> {
self.0.broadcaster.start(self.local());
let mut tasks = Vec::new();
if let Some(service) = state.get::<AuthService>() {
let hub = self.clone();
let state = state.clone();
let mut revocations = service.subscribe_revocations();
let mut role_changes = service.subscribe_role_changes();
let refresh_every = self.0.config.roles_refresh_secs;
tasks.push(tokio::spawn(async move {
let mut last = state.now();
let period = Duration::from_secs(refresh_every.max(1));
let mut refresh = tokio::time::interval_at(tokio::time::Instant::now() + period, period);
refresh.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
tokio::select! {
_ = state.shutdown().wait() => break,
received = revocations.recv() => match received {
Ok(revocation) => {
last = state.now();
hub.apply_revocation(&revocation);
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(missed)) => {
tracing::warn!(missed, "WebSocket: revocations lagged; re-reading the sessions table");
let since = UnixMillis(last.get().saturating_sub(LAG_OVERLAP_MS));
match service.revocations_since(&state, since).await {
Ok(list) => {
for (at, revocation) in list {
last = UnixMillis(last.get().max(at.get()));
hub.apply_revocation(&revocation);
}
}
Err(error) => tracing::warn!(%error, "WebSocket: reading revocations failed"),
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
},
changed = role_changes.recv() => match changed {
Ok(user) => hub.refresh_roles(&state, &service, vec![user]).await,
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
let users: Vec<UserId> = hub.read().users.keys().copied().collect();
hub.refresh_roles(&state, &service, users).await;
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
},
_ = refresh.tick(), if refresh_every > 0 => {
let users: Vec<UserId> = hub.read().users.keys().copied().collect();
hub.refresh_roles(&state, &service, users).await;
}
}
}
}));
}
let hub = self.clone();
let shutdown = state.shutdown().clone();
let grace = Duration::from_secs(state.config().server.shutdown_grace_secs);
tasks.push(tokio::spawn(async move {
shutdown.wait().await;
hub.begin_shutdown(grace);
}));
tasks
}
pub(crate) fn begin_shutdown(&self, grace: Duration) {
if self.0.closing.swap(true, Ordering::SeqCst) {
return;
}
*self.0.deadline.lock().unwrap_or_else(|p| p.into_inner()) = Some(Instant::now() + grace);
let registry = self.read();
let count = registry.conns.values().filter(|c| c.request_close(CloseCode::GOING_AWAY, Cow::Borrowed("the server is shutting down"))).count();
drop(registry);
if count > 0 {
tracing::info!(sockets = count, ">>> NBS: closing WebSockets (1001)");
}
}
pub(crate) async fn finish(&self, grace: Duration) {
self.begin_shutdown(grace);
let deadline = self.0.deadline.lock().unwrap_or_else(|p| p.into_inner()).unwrap_or_else(|| Instant::now() + grace);
while self.0.live.load(Ordering::SeqCst) > 0 && Instant::now() < deadline {
tokio::time::sleep(Duration::from_millis(20)).await;
}
let left = self.0.live.load(Ordering::SeqCst);
if left > 0 {
tracing::warn!(sockets = left, "shutdown grace period over; the remaining WebSockets are dropped");
self.0.kill.send_replace(true);
let until = Instant::now() + Duration::from_secs(2);
while self.0.live.load(Ordering::SeqCst) > 0 && Instant::now() < until {
tokio::time::sleep(Duration::from_millis(20)).await;
}
}
}
}
impl axum::extract::FromRef<AppState> for Hub {
fn from_ref(state: &AppState) -> Hub {
state.ws().clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn close_codes_are_sendable() {
for (code, sent) in [
(1000, 1000),
(1001, 1001),
(1004, 1011),
(1005, 1011),
(1006, 1011),
(1015, 1011),
(999, 1011),
(2999, 1011),
(4001, 4001),
(4999, 4999),
(5000, 1011),
] {
assert_eq!(sendable(CloseCode(code)).get(), sent, "{code}");
}
}
#[test]
fn deliveries_serialize() {
let delivery = Delivery::new(Target::Room(Arc::from("lobby")), r#"{"type":"x","data":1}"#);
let json = serde_json::to_string(&delivery).unwrap_or_default();
assert_eq!(json, r#"{"target":{"Room":"lobby"},"frame":"{\"type\":\"x\",\"data\":1}"}"#);
let back: Delivery = serde_json::from_str(&json).unwrap_or_else(|_| Delivery::new(Target::All, ""));
assert_eq!(back.target, delivery.target);
assert_eq!(back.frame.as_str(), delivery.frame.as_str());
assert_eq!(back.control, None);
let leave = Delivery::control(Target::User(UserId(7)), Control::LeaveRoom(Arc::from("chat:3")));
let json = serde_json::to_string(&leave).unwrap_or_default();
assert_eq!(json, r#"{"target":{"User":7},"frame":"","control":{"leave_room":"chat:3"}}"#);
let back: Delivery = serde_json::from_str(&json).unwrap_or_else(|_| Delivery::new(Target::All, ""));
assert_eq!(back.control, leave.control);
}
}