1use std::{str, sync::Arc, time::Duration};
2
3use axum::{
4 Error as AxumError,
5 extract::ws::{Message, WebSocket},
6};
7use futures_util::StreamExt;
8use o_sfu_protocol::wire::{ClientEnvelope, WebSocketCloseCode as CloseCode};
9use tokio::time::{Instant, sleep_until};
10use tokio_util::sync::CancellationToken;
11use tracing::{Instrument, Span, debug, field, info, info_span, instrument, warn};
12
13use super::{
14 WsReader, WsWriter,
15 admission::PreAuthWebSocketPermit,
16 controller::WebSocketServices,
17 handshake::{self, AuthenticatedJoin, HandshakeError, WebSocketAuth},
18 io::{close_writer_bounded, send_message_bounded, send_user_output_bounded},
19};
20use crate::{
21 application::user_session::{User, UserError, UserOutput},
22 config::UserConfig,
23 core::server::room::{
24 JoinUserRequest, RoomManagerJoinError, UserOutbound, UserOutboundEvent,
25 UserOutboundQueueLimits, UserOutboundReceiver, UserOutboundSender,
26 },
27 runtime::{
28 metrics::{RuntimeMetrics, WsSessionLoopExitReason as LoopExit},
29 telemetry::{
30 self,
31 schema::{event as telemetry_event, field as telemetry_field},
32 },
33 websocket_server::{ClientBatchDecodeFailureKind, decode_client_batch},
34 },
35};
36
37struct AuthenticatedSession {
38 _proof: WebSocketAuth,
39 writer: WsWriter,
40 reader: WsReader,
41 outbound: UserOutboundReceiver,
42 user: User,
43 user_config: UserConfig,
44 metrics: Arc<RuntimeMetrics>,
45 shutdown: CancellationToken,
46}
47
48enum SessionExit {
49 BeforeLoop(Option<CloseCode>),
50 Loop(LoopExit, Option<CloseCode>),
51}
52
53impl SessionExit {
54 const fn closing(reason: LoopExit, code: CloseCode) -> Self {
55 Self::Loop(reason, Some(code))
56 }
57}
58
59pub(super) async fn run(
60 socket: WebSocket,
61 services: WebSocketServices,
62 remote: Arc<str>,
63 permit: PreAuthWebSocketPermit,
64) {
65 async move {
66 Span::current().record(
67 telemetry_field::REMOTE_ADDRESS,
68 field::display(remote.as_ref()),
69 );
70 services.metrics.record_ws_connection_accepted();
71 let handshake_span = telemetry::ws_handshake_span();
72 handshake_span.record(
73 telemetry_field::REMOTE_ADDRESS,
74 field::display(remote.as_ref()),
75 );
76 if let Some(session) = establish(socket, services, remote, permit)
77 .instrument(handshake_span)
78 .await
79 {
80 session.serve().await;
81 }
82 }
83 .instrument(telemetry::ws_upgrade_span())
84 .await;
85}
86
87async fn establish(
88 mut socket: WebSocket,
89 services: WebSocketServices,
90 remote: Arc<str>,
91 permit: PreAuthWebSocketPermit,
92) -> Option<AuthenticatedSession> {
93 let _guard = services.metrics.track_ws_handshake();
94 let auth = {
95 let _guard = services.metrics.track_ws_authentication();
96 handshake::authenticate(&services, &mut socket, remote.as_ref()).await
97 };
98 let (mut writer, reader) = socket.split();
99 let join = match auth {
100 Ok(join) => join,
101 Err(HandshakeError::PeerClosed) => return None,
102 Err(HandshakeError::Rejected(code)) => {
103 handshake::reject(
104 &services,
105 &mut writer,
106 code,
107 remote.as_ref(),
108 "rejecting websocket during authentication",
109 )
110 .await;
111 return None;
112 }
113 Err(HandshakeError::Shutdown) => {
114 close_writer_bounded(&mut writer, CloseCode::Leaving).await;
115 return None;
116 }
117 };
118 drop(permit);
119 if services.shutdown.is_cancelled() {
120 close_writer_bounded(&mut writer, CloseCode::Leaving).await;
121 return None;
122 }
123 let (proof, user, outbound) = admit(&services, join, remote, &mut writer).await?;
124 let mut session = AuthenticatedSession {
125 _proof: proof,
126 writer,
127 reader,
128 outbound,
129 user,
130 user_config: services.user,
131 metrics: Arc::clone(&services.metrics),
132 shutdown: services.shutdown,
133 };
134 session.metrics.record_ws_user_joined();
135 session.record_current_span();
136 if session.shutdown.is_cancelled() {
137 session
138 .finish(SessionExit::BeforeLoop(Some(CloseCode::Leaving)))
139 .await;
140 return None;
141 }
142 session.start().await
143}
144
145#[instrument(
146 name = "room.join",
147 skip_all,
148 fields(room_id = %room.uuid(), user_id = ?claims.user_id)
149)]
150async fn admit(
151 services: &WebSocketServices,
152 AuthenticatedJoin {
153 room,
154 claims,
155 proof,
156 }: AuthenticatedJoin,
157 remote: Arc<str>,
158 writer: &mut WsWriter,
159) -> Option<(WebSocketAuth, User, UserOutboundReceiver)> {
160 let user_id = claims.user_id;
161 let limits = UserOutboundQueueLimits::new(
162 services.user.outbound_queue_capacity,
163 services.user.outbound_queue_byte_capacity,
164 );
165 let (outbound_tx, outbound) =
166 UserOutboundSender::channel_with_limits(limits, Arc::clone(&services.metrics));
167 match services
168 .sfu_core
169 .admit_user(
170 room.uuid(),
171 JoinUserRequest {
172 user_id: user_id.clone(),
173 label: claims.label,
174 permissions: claims.permissions.unwrap_or_default(),
175 sender: outbound_tx,
176 },
177 )
178 .await
179 {
180 Ok(session) => Some((proof, User::new(session, remote), outbound)),
181 Err(_error) if services.shutdown.is_cancelled() => {
182 close_writer_bounded(writer, CloseCode::Leaving).await;
183 None
184 }
185 Err(error) => {
186 let code = match error {
187 RoomManagerJoinError::RoomFull => CloseCode::RoomFull,
188 RoomManagerJoinError::MissingRoom | RoomManagerJoinError::RouterState => {
189 CloseCode::AuthFailed
190 }
191 };
192 warn!(
193 event = telemetry_event::WS_JOIN_FAILED,
194 ?user_id,
195 remote_address = remote.as_ref(),
196 ?error,
197 close_code = u16::from(code),
198 "rejecting websocket because the authenticated user could not join the room"
199 );
200 handshake::reject(
201 services,
202 writer,
203 code,
204 remote.as_ref(),
205 "rejecting websocket during user join",
206 )
207 .await;
208 None
209 }
210 }
211}
212
213impl AuthenticatedSession {
214 async fn serve(mut self) {
215 self.record_current_span();
216 self.metrics.record_ws_user_loop_started();
217 let exit = self.run_loop().await;
218 self.finish(exit).await;
219 }
220
221 fn record_current_span(&self) {
222 let span = Span::current();
223 span.record("room_id", field::display(self.user.room_id()));
224 span.record("user_id", field::debug(self.user.user_id()));
225 span.record("connection_id", field::debug(self.user.connection_id()));
226 span.record(
227 telemetry_field::REMOTE_ADDRESS,
228 field::display(self.user.remote_address()),
229 );
230 }
231
232 async fn start(mut self) -> Option<Self> {
233 let metrics = Arc::clone(&self.metrics);
234 let _guard = metrics.track_ws_user_initialization();
235 let span = telemetry::activated_span(info_span!(
236 "user.initialize",
237 room_id = %self.user.room_id(),
238 user_id = ?self.user.user_id(),
239 connection_id = ?self.user.connection_id(),
240 remote_address = %self.user.remote_address()
241 ));
242 async move {
243 match self.start_inner().await {
244 Ok(()) => Some(self),
245 Err(exit) => {
246 self.finish(exit).await;
247 None
248 }
249 }
250 }
251 .instrument(span)
252 .await
253 }
254
255 async fn start_inner(&mut self) -> Result<(), SessionExit> {
256 let output = self.user.start().await;
257 if self.shutdown.is_cancelled() {
258 return Err(SessionExit::BeforeLoop(Some(CloseCode::Leaving)));
259 }
260 let output = match output {
261 Ok(output) => output,
262 Err(_error) => {
263 warn!(
264 event = telemetry_event::WS_JOIN_FAILED,
265 user_id = ?self.user.user_id(),
266 connection_id = ?self.user.connection_id(),
267 remote_address = self.user.remote_address(),
268 outcome = "user_initialize_failed",
269 "failed to initialize websocket user"
270 );
271 self.metrics.record_ws_user_initialize_failure();
272 return Err(SessionExit::BeforeLoop(None));
273 }
274 };
275 let sent = send_user_output_bounded(&mut self.writer, output).await;
276 if self.shutdown.is_cancelled() {
277 return Err(SessionExit::BeforeLoop(Some(CloseCode::Leaving)));
278 }
279 if sent.is_ok() {
280 return Ok(());
281 }
282 debug!(
283 user_id = ?self.user.user_id(),
284 connection_id = ?self.user.connection_id(),
285 "failed to send user startup payload"
286 );
287 self.metrics.record_ws_startup_send_failure();
288 warn!(
289 event = telemetry_event::WS_JOIN_FAILED,
290 user_id = ?self.user.user_id(),
291 connection_id = ?self.user.connection_id(),
292 remote_address = self.user.remote_address(),
293 outcome = "startup_send_failed",
294 "failed to send websocket user startup payload"
295 );
296 Err(SessionExit::BeforeLoop(None))
297 }
298
299 async fn finish(&mut self, exit: SessionExit) {
300 let (reason, close) = match exit {
301 SessionExit::BeforeLoop(close) => (None, close),
302 SessionExit::Loop(reason, close) => (Some(reason), close),
303 };
304 if let Some(close) = close {
305 close_writer_bounded(&mut self.writer, close).await;
306 }
307 if let Some(reason) = reason {
308 self.metrics.record_ws_user_loop_exit(reason);
309 info!(
310 event = telemetry_event::WS_CONNECTION_CLOSED,
311 connection_id = ?self.user.connection_id(),
312 remote_address = self.user.remote_address(),
313 ?reason,
314 "closing websocket user"
315 );
316 }
317 self.user.close().await;
318 }
319
320 fn shutdown_exit(&self) -> Option<SessionExit> {
321 self.shutdown.is_cancelled().then_some(SessionExit::closing(
322 LoopExit::RuntimeShutdown,
323 CloseCode::Leaving,
324 ))
325 }
326
327 #[allow(
329 clippy::cognitive_complexity,
330 reason = "all session wake sources stay in one owner loop"
331 )]
332 async fn run_loop(&mut self) -> SessionExit {
333 let ping_interval = Duration::from_millis(self.user_config.ping_interval_ms);
334 let ping_timeout = Duration::from_millis(self.user_config.timeout_ms);
335 let mut next_ping_at = Instant::now() + ping_interval;
336 let mut next_health_at = next_ping_at;
337 let mut pong = None;
338 let shutdown = self.shutdown.clone();
339 loop {
340 let health_tick = sleep_until(next_health_at);
341 tokio::pin!(health_tick);
342 let ping_tick = sleep_until(next_ping_at);
343 tokio::pin!(ping_tick);
344 let pong_deadline = pong;
345 tokio::select! {
346 biased;
347 () = shutdown.cancelled() => {
348 return SessionExit::closing(LoopExit::RuntimeShutdown, CloseCode::Leaving);
349 }
350 () = &mut health_tick => {
351 next_health_at = Instant::now() + ping_interval;
352 if let Some(exit) = self.check_transport() {
353 return exit;
354 }
355 }
356 () = &mut ping_tick, if pong.is_none() => {
357 if let Some(exit) = self.check_transport() {
358 return exit;
359 }
360 if send_message_bounded(&mut self.writer, Message::Ping(Vec::new().into()))
361 .await
362 .is_err()
363 {
364 debug!("failed to send websocket ping frame");
365 return SessionExit::Loop(LoopExit::OutboundMessageSendFailure, None);
366 }
367 let now = Instant::now();
368 next_ping_at = now + ping_interval;
369 pong = Some(now + ping_timeout);
370 }
371 () = async {
372 if let Some(deadline) = pong_deadline {
373 sleep_until(deadline).await;
374 }
375 }, if pong_deadline.is_some() => {
376 debug!("timed out waiting for websocket pong");
377 return SessionExit::closing(LoopExit::PingTimeout, CloseCode::Error);
378 }
379 outbound = self.outbound.recv_event() => {
380 if let Some(exit) = self.handle_outbound_event(outbound).await {
381 return exit;
382 }
383 }
384 message = self.reader.next() => {
385 if let Some(exit) = self.handle_socket_event(message, &mut pong).await {
386 return exit;
387 }
388 }
389 }
390 }
391 }
392
393 fn check_transport(&self) -> Option<SessionExit> {
394 if !self.user.transport_disconnected() {
395 return None;
396 }
397 debug!("closing websocket because the underlying RTC transport disconnected");
398 Some(SessionExit::closing(
399 LoopExit::TransportDisconnected,
400 CloseCode::Error,
401 ))
402 }
403
404 async fn handle_socket_event(
405 &mut self,
406 message: Option<Result<Message, AxumError>>,
407 pong: &mut Option<Instant>,
408 ) -> Option<SessionExit> {
409 let message = match message {
410 Some(Ok(message)) => message,
411 Some(Err(_error)) => {
412 debug!("websocket reader returned an error");
413 return Some(SessionExit::Loop(LoopExit::ReaderError, None));
414 }
415 None => {
416 debug!("websocket user closed the socket");
417 return Some(SessionExit::Loop(LoopExit::UserClosed, None));
418 }
419 };
420 self.handle_frame(message, pong).await
421 }
422
423 async fn handle_frame(
424 &mut self,
425 message: Message,
426 pong: &mut Option<Instant>,
427 ) -> Option<SessionExit> {
428 match message {
429 Message::Ping(payload) => {
430 if send_message_bounded(&mut self.writer, Message::Pong(payload))
431 .await
432 .is_err()
433 {
434 debug!("failed to send websocket pong frame");
435 return Some(SessionExit::Loop(
436 LoopExit::OutboundMessageSendFailure,
437 None,
438 ));
439 }
440 None
441 }
442 Message::Pong(_) => {
443 *pong = None;
444 None
445 }
446 Message::Close(frame) => {
447 debug!(?frame, "websocket user sent close frame");
448 Some(SessionExit::Loop(LoopExit::BusBreak, None))
449 }
450 Message::Text(payload) => self.handle_text(&payload).await,
451 Message::Binary(payload) => self.handle_binary(&payload).await,
452 }
453 }
454
455 async fn handle_binary(&mut self, payload: &[u8]) -> Option<SessionExit> {
456 let Ok(payload) = str::from_utf8(payload) else {
457 self.metrics.record_ws_bus_invalid_input_failure();
458 warn!("received websocket binary frame with invalid UTF-8");
459 return Some(self.client_error(CloseCode::ProtocolError).await);
460 };
461 self.handle_text(payload).await
462 }
463
464 async fn handle_text(&mut self, payload: &str) -> Option<SessionExit> {
465 let batch = match decode_client_batch(payload) {
466 Ok(batch) => batch,
467 Err(error) => {
468 let failure = error.kind();
469 match failure {
470 ClientBatchDecodeFailureKind::InvalidInput => {
471 self.metrics.record_ws_bus_invalid_input_failure();
472 }
473 ClientBatchDecodeFailureKind::UnsupportedFeature => {
474 self.metrics.record_ws_bus_unsupported_feature_failure();
475 }
476 }
477 warn!(?failure, "failed to decode client websocket batch");
478 return Some(self.client_error(CloseCode::ProtocolError).await);
479 }
480 };
481 self.metrics.record_ws_bus_batch_received(batch.len());
482 let mut output = UserOutput::new();
483 for envelope in batch {
484 if let Some(exit) = self.shutdown_exit() {
485 return Some(exit);
486 }
487 match &envelope {
488 ClientEnvelope::Request { .. } => self.metrics.record_ws_bus_client_request(),
489 ClientEnvelope::Message(_) => self.metrics.record_ws_bus_client_message(),
490 ClientEnvelope::Response { .. } => {}
491 }
492 let result = self.user.apply_client_envelope(envelope).await;
493 if let Some(exit) = self.shutdown_exit() {
494 return Some(exit);
495 }
496 match result {
497 Ok(user_output) => output.extend(user_output),
498 Err(error) => return Some(self.client_error(map_user_error(error)).await),
499 }
500 }
501 let result = send_user_output_bounded(&mut self.writer, output).await;
502 if let Some(exit) = self.shutdown_exit() {
503 return Some(exit);
504 }
505 match result {
506 Ok(_sent) => None,
507 Err(code) => Some(SessionExit::closing(LoopExit::BusBreak, code)),
508 }
509 }
510
511 async fn client_error(&self, fallback_code: CloseCode) -> SessionExit {
512 let close_code = if self.user.is_current_connection().await {
513 fallback_code
514 } else {
515 CloseCode::Kicked
516 };
517 self.shutdown_exit()
518 .unwrap_or(SessionExit::closing(LoopExit::BusBreak, close_code))
519 }
520
521 async fn handle_outbound_event(&mut self, outbound: UserOutboundEvent) -> Option<SessionExit> {
522 match outbound {
523 UserOutboundEvent::Message(UserOutbound::Close(_)) => {
524 Some(self.outbound_error(CloseCode::Kicked, false))
525 }
526 UserOutboundEvent::Message(outbound) => {
527 let result = self.user.apply_room_outbound(outbound).await;
528 if let Some(exit) = self.shutdown_exit() {
529 return Some(exit);
530 }
531 let output = match result {
532 Ok(output) => output,
533 Err(error) => {
534 return Some(self.outbound_error(map_user_error(error), false));
535 }
536 };
537 let envelope_count = output.len();
538 let result = send_user_output_bounded(&mut self.writer, output).await;
539 if let Some(exit) = self.shutdown_exit() {
540 return Some(exit);
541 }
542 match result {
543 Ok(batch_count) => {
544 self.metrics
545 .record_ws_bus_batches_sent(batch_count, envelope_count);
546 None
547 }
548 Err(code) => Some(self.outbound_error(code, true)),
549 }
550 }
551 UserOutboundEvent::Overflow(overflow) => {
552 warn!(
553 capacity = overflow.capacity(),
554 byte_capacity = overflow.byte_capacity(),
555 queued_bytes = overflow.queued_bytes(),
556 message_bytes = overflow.message_bytes(),
557 overflow_kind = ?overflow.kind(),
558 "closing websocket because the outbound queue overflowed"
559 );
560 Some(SessionExit::closing(
561 LoopExit::OutboundQueueOverflow,
562 CloseCode::Kicked,
563 ))
564 }
565 UserOutboundEvent::Closed => {
566 debug!("user outbound room closed");
567 Some(SessionExit::Loop(LoopExit::OutboundChannelClosed, None))
568 }
569 }
570 }
571
572 fn outbound_error(&self, code: CloseCode, log_send_failure: bool) -> SessionExit {
573 if code == CloseCode::Kicked {
574 debug!(
575 close_code = u16::from(code),
576 "closing websocket from outbound signal"
577 );
578 return SessionExit::closing(LoopExit::OutboundCloseSignal, CloseCode::Kicked);
579 }
580 self.metrics.record_ws_bus_send_failure();
581 if log_send_failure {
582 debug!(
583 close_code = u16::from(code),
584 "failed to send outbound user event"
585 );
586 }
587 let close = matches!(
588 code,
589 CloseCode::Clean
590 | CloseCode::Leaving
591 | CloseCode::RoomFull
592 | CloseCode::AuthFailed
593 | CloseCode::AuthTimeout
594 )
595 .then_some(code);
596 SessionExit::Loop(LoopExit::OutboundMessageSendFailure, close)
597 }
598}
599
600fn map_user_error(error: UserError) -> CloseCode {
601 match error {
602 UserError::ProtocolViolation => CloseCode::ProtocolError,
603 UserError::Kicked => CloseCode::Kicked,
604 UserError::InternalError => CloseCode::Error,
605 }
606}