Skip to main content

o_sfu/runtime/websocket_server/
controller.rs

1//! websocket controller for one upgraded socket
2//!
3//! this module bounds upgrade admission before handing the socket to
4//! [`super::session::run`]
5
6use std::{net::IpAddr, sync::Arc};
7
8use axum::{
9    extract::{FromRef, State, ws::WebSocketUpgrade},
10    http::StatusCode,
11    response::{IntoResponse, Response},
12};
13use tokio_util::{sync::CancellationToken, task::TaskTracker};
14use tracing::warn;
15
16use super::{
17    admission::{PreAuthWebSocketAdmissionRejection, admit_rejection_log},
18    io::MAX_CLIENT_FRAME_BYTES,
19    session,
20};
21use crate::{
22    config::{DeadlineDuration, UserConfig},
23    core::prelude::SfuCore,
24    runtime::{
25        RuntimeMetrics, RuntimeState,
26        request_origin::RequestOrigin,
27        room::RoomManager,
28        telemetry::{metrics::WsPreAuthRejection, schema::event as telemetry_event},
29    },
30};
31
32pub(crate) struct WebSocketServices {
33    pub(super) authentication_timeout: DeadlineDuration,
34    max_pre_auth_websocket_sessions: usize,
35    max_pre_auth_websocket_sessions_per_origin: usize,
36    pub(super) user: UserConfig,
37    pub(super) room_manager: Arc<RoomManager>,
38    pub(super) sfu_core: SfuCore,
39    pub(super) metrics: Arc<RuntimeMetrics>,
40    pub(super) shutdown: CancellationToken,
41    sessions: TaskTracker,
42    pre_auth_websocket_admission: super::PreAuthWebSocketAdmission,
43}
44
45impl FromRef<RuntimeState> for WebSocketServices {
46    fn from_ref(state: &RuntimeState) -> Self {
47        Self {
48            authentication_timeout: state.config.auth.authentication_timeout,
49            max_pre_auth_websocket_sessions: state.config.auth.max_pre_auth_websocket_sessions,
50            max_pre_auth_websocket_sessions_per_origin: state
51                .config
52                .auth
53                .max_pre_auth_websocket_sessions_per_origin,
54            user: state.config.user,
55            room_manager: Arc::clone(&state.room_manager),
56            sfu_core: state.sfu_core.clone(),
57            metrics: Arc::clone(&state.metrics),
58            shutdown: state.session_shutdown.clone(),
59            sessions: state.session_tasks.clone(),
60            pre_auth_websocket_admission: state.pre_auth_websocket_admission.clone(),
61        }
62    }
63}
64
65pub(crate) async fn upgrade(
66    State(services): State<WebSocketServices>,
67    origin: RequestOrigin,
68    websocket: WebSocketUpgrade,
69) -> Response {
70    let pre_auth_permit = match services
71        .pre_auth_websocket_admission
72        .try_acquire(origin.remote_address)
73    {
74        Ok(permit) => permit,
75        Err(rejection) => {
76            reject_pre_auth_admission(&services, origin.remote_address, rejection);
77            return StatusCode::SERVICE_UNAVAILABLE.into_response();
78        }
79    };
80    let remote_address = Arc::<str>::from(format_remote_address(origin.remote_address));
81    let session_task = services.sessions.token();
82    websocket
83        .max_message_size(MAX_CLIENT_FRAME_BYTES)
84        .max_frame_size(MAX_CLIENT_FRAME_BYTES)
85        .on_upgrade(move |socket| async move {
86            session::run(socket, services, remote_address, pre_auth_permit).await;
87            drop(session_task);
88        })
89}
90
91fn reject_pre_auth_admission(
92    services: &WebSocketServices,
93    remote_address: Option<IpAddr>,
94    rejection: PreAuthWebSocketAdmissionRejection,
95) {
96    services
97        .metrics
98        .record_ws_pre_auth_rejection(match rejection {
99            PreAuthWebSocketAdmissionRejection::Global => WsPreAuthRejection::Global,
100            PreAuthWebSocketAdmissionRejection::Origin => WsPreAuthRejection::Origin,
101        });
102    let Some(suppressed_rejections) = admit_rejection_log() else {
103        return;
104    };
105    let remote_address = format_remote_address(remote_address);
106    match rejection {
107        PreAuthWebSocketAdmissionRejection::Global => {
108            warn!(
109                event = telemetry_event::WS_HANDSHAKE_REJECTED,
110                remote_address,
111                suppressed_rejections,
112                max_pre_auth_websocket_sessions = services.max_pre_auth_websocket_sessions,
113                "rejecting websocket upgrade because global pre-auth admission is full"
114            );
115        }
116        PreAuthWebSocketAdmissionRejection::Origin => {
117            warn!(
118                event = telemetry_event::WS_HANDSHAKE_REJECTED,
119                remote_address,
120                suppressed_rejections,
121                max_pre_auth_websocket_sessions_per_origin =
122                    services.max_pre_auth_websocket_sessions_per_origin,
123                "rejecting websocket upgrade because origin pre-auth admission is full"
124            );
125        }
126    }
127}
128
129fn format_remote_address(address: Option<IpAddr>) -> String {
130    address.map_or_else(|| "unknown".to_owned(), |address| address.to_string())
131}