o_sfu/runtime/websocket_server/
controller.rs1use 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}