use std::time::Duration;
use crate::axum::body::Bytes;
use crate::axum::extract::ws::{self, CloseFrame, Message, WebSocket, WebSocketUpgrade};
use crate::axum::http::{HeaderMap, StatusCode};
use crate::axum::response::{IntoResponse, Response};
use crate::realtime::channel::Broadcast;
use crate::realtime::error::{ChannelError, ProtocolHint, RealtimeError, admission_status};
use crate::realtime::origin::{OriginDecision, OriginPolicy};
use crate::realtime::registry::Registry;
use crate::realtime::shutdown::ShutdownConfig;
pub trait Authorizer: Clone + Send + Sync + 'static {
fn authorize(
&self,
headers: &HeaderMap,
channel_id: &str,
) -> impl std::future::Future<Output = Option<Broadcast>> + Send;
}
#[derive(Clone)]
pub struct AllowAll {
broadcast: Broadcast,
}
impl AllowAll {
#[must_use]
pub fn new(broadcast: Broadcast) -> Self {
Self { broadcast }
}
}
impl Authorizer for AllowAll {
async fn authorize(&self, _headers: &HeaderMap, _channel_id: &str) -> Option<Broadcast> {
Some(self.broadcast.clone())
}
}
#[derive(Clone, Copy, Debug)]
pub struct WsLimits {
pub max_message_size: usize,
pub max_frame_size: usize,
pub heartbeat_interval: Duration,
pub heartbeat_timeout: Duration,
}
impl WsLimits {
#[must_use]
pub fn conservative() -> Self {
Self {
max_message_size: 64 * 1024,
max_frame_size: 64 * 1024,
heartbeat_interval: Duration::from_secs(20),
heartbeat_timeout: Duration::from_secs(40),
}
}
}
#[derive(Clone)]
pub struct WebSocketEndpoint<A> {
authorizer: A,
origin: OriginPolicy,
registry: Registry,
limits: WsLimits,
shutdown: ShutdownConfig,
}
impl<A: Authorizer> WebSocketEndpoint<A> {
#[must_use]
pub fn new(
authorizer: A,
origin: OriginPolicy,
registry: Registry,
limits: WsLimits,
shutdown: ShutdownConfig,
) -> Self {
Self {
authorizer,
origin,
registry,
limits,
shutdown,
}
}
pub async fn handle(
self,
ws: WebSocketUpgrade,
headers: HeaderMap,
channel_id: String,
) -> Response {
if self.origin.authorize(headers.get("origin")) == OriginDecision::Denied {
tracing::debug!(
target: "arcature::realtime::ws::admit",
realtime_transport = "ws",
error_category = "origin",
"realtime upgrade rejected: origin policy"
);
return admission_status(&RealtimeError::Origin).into_response();
}
let broadcast = match self.authorizer.authorize(&headers, &channel_id).await {
Some(bc) => bc,
None => {
tracing::debug!(
target: "arcature::realtime::ws::admit",
realtime_transport = "ws",
error_category = "authz",
"realtime upgrade rejected: authorization"
);
return admission_status(&RealtimeError::Unauthorized).into_response();
}
};
let guard = match self.registry.acquire(self.shutdown.max_connections()) {
Ok(g) => g,
Err(_) => {
tracing::debug!(
target: "arcature::realtime::ws::admit",
realtime_transport = "ws",
error_category = "limit",
"realtime upgrade rejected: connection limit"
);
return admission_status(&RealtimeError::ConnectionLimit).into_response();
}
};
let limits = self.limits;
let shutdown = self.shutdown;
ws.max_message_size(limits.max_message_size)
.max_frame_size(limits.max_frame_size)
.on_upgrade(move |socket| run_connection(socket, broadcast, guard, limits, shutdown))
}
}
async fn run_connection(
mut socket: WebSocket,
broadcast: Broadcast,
guard: crate::realtime::registry::ConnectionGuard,
limits: WsLimits,
shutdown: ShutdownConfig,
) {
let mut sub = broadcast.subscribe();
let span = tracing::debug_span!(
target: "arcature::realtime::ws",
"ws.connection",
realtime_transport = "ws",
);
let _enter = span.enter();
let mut last_seen = tokio::time::Instant::now();
let mut heartbeat: Option<tokio::time::Interval> = if limits.heartbeat_interval.is_zero() {
None
} else {
let mut i = tokio::time::interval(limits.heartbeat_interval);
i.tick().await;
Some(i)
};
loop {
if shutdown.is_draining() {
break;
}
tokio::select! {
_ = shutdown.drain_notified() => {
continue;
}
msg = socket.recv() => {
match msg {
Some(Ok(Message::Text(_))) | Some(Ok(Message::Binary(_))) => {
last_seen = tokio::time::Instant::now();
}
Some(Ok(Message::Pong(_))) | Some(Ok(Message::Ping(_))) => {
last_seen = tokio::time::Instant::now();
}
Some(Ok(Message::Close(_))) => break,
Some(Err(_)) => {
close_protocol(&mut socket, ProtocolHint::Stream).await;
break;
}
None => break,
}
}
recv = sub.recv(), if !shutdown.is_draining() => {
match recv {
Ok(payload) => {
forward_payload(&mut socket, payload).await;
}
Err(ChannelError::Lagged) => {
close_protocol(&mut socket, ProtocolHint::Malformed).await;
break;
}
Err(ChannelError::Closed) | Err(ChannelError::Full) => break,
}
}
_ = ping_tick(heartbeat.as_mut()), if heartbeat.is_some() => {
let now = tokio::time::Instant::now();
if now.duration_since(last_seen) > limits.heartbeat_timeout {
let _ = socket
.send(Message::Close(Some(CloseFrame {
code: ws::close_code::AGAIN,
reason: "heartbeat timeout".into(),
})))
.await;
break;
}
let _ = socket.send(Message::Ping(Bytes::new())).await;
}
}
}
let _ = socket
.send(Message::Close(Some(CloseFrame {
code: ws::close_code::AWAY,
reason: "server draining".into(),
})))
.await;
drop(guard);
}
async fn forward_payload(
socket: &mut WebSocket,
payload: crate::realtime::channel::ChannelPayload,
) {
match std::str::from_utf8(payload.as_bytes()) {
Ok(s) => {
let _ = socket.send(Message::text(s.to_string())).await;
}
Err(_) => {
let _ = socket
.send(Message::binary(payload.as_bytes().to_vec()))
.await;
}
}
}
async fn ping_tick(tick: Option<&mut tokio::time::Interval>) {
if let Some(t) = tick {
t.tick().await;
}
}
async fn close_protocol(socket: &mut WebSocket, hint: ProtocolHint) {
let _ = socket
.send(Message::Close(Some(CloseFrame {
code: ws::close_code::ERROR,
reason: hint.to_string().into(),
})))
.await;
}
const _: () = {
fn _assert(status: StatusCode) -> Response {
status.into_response()
}
};