1use std::{str, sync::Arc};
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).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::Shutdown) => {
103 close_writer_bounded(&mut writer, CloseCode::Leaving).await;
104 return None;
105 }
106 Err(error) => {
107 handshake::reject(
108 &services,
109 &mut writer,
110 error.close_code(),
111 remote.as_ref(),
112 error,
113 )
114 .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.path_segment())
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 => CloseCode::AuthFailed,
189 RoomManagerJoinError::NoUsableWorker | RoomManagerJoinError::RouterState => {
190 CloseCode::Error
191 }
192 };
193 warn!(
194 event = telemetry_event::WS_JOIN_FAILED,
195 ?user_id,
196 remote_address = remote.as_ref(),
197 ?error,
198 close_code = u16::from(code),
199 "rejecting websocket because the authenticated user could not join the room"
200 );
201 handshake::reject(
202 services,
203 writer,
204 code,
205 remote.as_ref(),
206 "rejecting websocket during user join",
207 )
208 .await;
209 None
210 }
211 }
212}
213
214impl AuthenticatedSession {
215 async fn serve(mut self) {
216 self.record_current_span();
217 self.metrics.record_ws_user_loop_started();
218 let exit = self.run_loop().await;
219 self.finish(exit).await;
220 }
221
222 fn record_current_span(&self) {
223 let span = Span::current();
224 span.record("room_id", field::display(self.user.room_id()));
225 span.record(
226 "user_id",
227 field::display(self.user.user_id().path_segment()),
228 );
229 span.record("connection_id", self.user.connection_id().as_u64());
230 span.record(
231 telemetry_field::REMOTE_ADDRESS,
232 field::display(self.user.remote_address()),
233 );
234 }
235
236 async fn start(mut self) -> Option<Self> {
237 let metrics = Arc::clone(&self.metrics);
238 let _guard = metrics.track_ws_user_initialization();
239 let span = telemetry::activated_span(info_span!(
240 "user.initialize",
241 room_id = %self.user.room_id(),
242 user_id = %self.user.user_id().path_segment(),
243 connection_id = self.user.connection_id().as_u64(),
244 remote_address = %self.user.remote_address()
245 ));
246 async move {
247 match self.start_inner().await {
248 Ok(()) => Some(self),
249 Err(exit) => {
250 self.finish(exit).await;
251 None
252 }
253 }
254 }
255 .instrument(span)
256 .await
257 }
258
259 async fn start_inner(&mut self) -> Result<(), SessionExit> {
260 let output = self.user.start().await;
261 if self.shutdown.is_cancelled() {
262 return Err(SessionExit::BeforeLoop(Some(CloseCode::Leaving)));
263 }
264 let output = match output {
265 Ok(output) => output,
266 Err(_error) => {
267 warn!(
268 event = telemetry_event::WS_JOIN_FAILED,
269 user_id = ?self.user.user_id(),
270 connection_id = ?self.user.connection_id(),
271 remote_address = self.user.remote_address(),
272 outcome = "user_initialize_failed",
273 "failed to initialize websocket user"
274 );
275 self.metrics.record_ws_user_initialize_failure();
276 return Err(SessionExit::BeforeLoop(None));
277 }
278 };
279 let sent = send_user_output_bounded(&mut self.writer, output).await;
280 if self.shutdown.is_cancelled() {
281 return Err(SessionExit::BeforeLoop(Some(CloseCode::Leaving)));
282 }
283 if sent.is_ok() {
284 return Ok(());
285 }
286 debug!(
287 user_id = ?self.user.user_id(),
288 connection_id = ?self.user.connection_id(),
289 "failed to send user startup payload"
290 );
291 self.metrics.record_ws_startup_send_failure();
292 warn!(
293 event = telemetry_event::WS_JOIN_FAILED,
294 user_id = ?self.user.user_id(),
295 connection_id = ?self.user.connection_id(),
296 remote_address = self.user.remote_address(),
297 outcome = "startup_send_failed",
298 "failed to send websocket user startup payload"
299 );
300 Err(SessionExit::BeforeLoop(None))
301 }
302
303 async fn finish(&mut self, exit: SessionExit) {
304 let (reason, close) = match exit {
305 SessionExit::BeforeLoop(close) => (None, close),
306 SessionExit::Loop(reason, close) => (Some(reason), close),
307 };
308 if let Some(close) = close {
309 close_writer_bounded(&mut self.writer, close).await;
310 }
311 if let Some(reason) = reason {
312 self.metrics.record_ws_user_loop_exit(reason);
313 info!(
314 event = telemetry_event::WS_CONNECTION_CLOSED,
315 connection_id = ?self.user.connection_id(),
316 remote_address = self.user.remote_address(),
317 ?reason,
318 "closing websocket user"
319 );
320 }
321 self.user.close().await;
322 }
323
324 fn shutdown_exit(&self) -> Option<SessionExit> {
325 self.shutdown.is_cancelled().then_some(SessionExit::closing(
326 LoopExit::RuntimeShutdown,
327 CloseCode::Leaving,
328 ))
329 }
330
331 #[expect(
333 clippy::cognitive_complexity,
334 reason = "all session wake sources stay in one owner loop"
335 )]
336 async fn run_loop(&mut self) -> SessionExit {
337 let ping_interval = self.user_config.ping_interval.as_duration();
338 let ping_timeout = self.user_config.timeout.as_duration();
339 let mut next_ping_at = Instant::now() + ping_interval;
340 let mut next_health_at = next_ping_at;
341 let mut pong = None;
342 let shutdown = self.shutdown.clone();
343 loop {
344 let health_tick = sleep_until(next_health_at);
345 tokio::pin!(health_tick);
346 let ping_tick = sleep_until(next_ping_at);
347 tokio::pin!(ping_tick);
348 let pong_deadline = pong;
349 tokio::select! {
350 biased;
351 () = shutdown.cancelled() => {
352 return SessionExit::closing(LoopExit::RuntimeShutdown, CloseCode::Leaving);
353 }
354 () = &mut health_tick => {
355 next_health_at = Instant::now() + ping_interval;
356 if let Some(exit) = self.check_transport() {
357 return exit;
358 }
359 }
360 () = &mut ping_tick, if pong.is_none() => {
361 if let Some(exit) = self.check_transport() {
362 return exit;
363 }
364 if send_message_bounded(&mut self.writer, Message::Ping(Vec::new().into()))
365 .await
366 .is_err()
367 {
368 debug!("failed to send websocket ping frame");
369 return SessionExit::Loop(LoopExit::OutboundMessageSendFailure, None);
370 }
371 let now = Instant::now();
372 next_ping_at = now + ping_interval;
373 pong = Some(now + ping_timeout);
374 }
375 () = async {
376 if let Some(deadline) = pong_deadline {
377 sleep_until(deadline).await;
378 }
379 }, if pong_deadline.is_some() => {
380 debug!("timed out waiting for websocket pong");
381 return SessionExit::closing(LoopExit::PingTimeout, CloseCode::Error);
382 }
383 outbound = self.outbound.recv_event() => {
384 if let Some(exit) = self.handle_outbound_event(outbound).await {
385 return exit;
386 }
387 }
388 message = self.reader.next() => {
389 if let Some(exit) = self.handle_socket_event(message, &mut pong).await {
390 return exit;
391 }
392 }
393 }
394 }
395 }
396
397 fn check_transport(&self) -> Option<SessionExit> {
398 if !self.user.transport_disconnected() {
399 return None;
400 }
401 debug!("closing websocket because the underlying RTC transport disconnected");
402 Some(SessionExit::closing(
403 LoopExit::TransportDisconnected,
404 CloseCode::Error,
405 ))
406 }
407
408 async fn handle_socket_event(
409 &mut self,
410 message: Option<Result<Message, AxumError>>,
411 pong: &mut Option<Instant>,
412 ) -> Option<SessionExit> {
413 let message = match message {
414 Some(Ok(message)) => message,
415 Some(Err(_error)) => {
416 debug!("websocket reader returned an error");
417 return Some(SessionExit::Loop(LoopExit::ReaderError, None));
418 }
419 None => {
420 debug!("websocket user closed the socket");
421 return Some(SessionExit::Loop(LoopExit::UserClosed, None));
422 }
423 };
424 self.handle_frame(message, pong).await
425 }
426
427 async fn handle_frame(
428 &mut self,
429 message: Message,
430 pong: &mut Option<Instant>,
431 ) -> Option<SessionExit> {
432 match message {
433 Message::Ping(payload) => {
434 if send_message_bounded(&mut self.writer, Message::Pong(payload))
435 .await
436 .is_err()
437 {
438 debug!("failed to send websocket pong frame");
439 return Some(SessionExit::Loop(
440 LoopExit::OutboundMessageSendFailure,
441 None,
442 ));
443 }
444 None
445 }
446 Message::Pong(_) => {
447 *pong = None;
448 None
449 }
450 Message::Close(frame) => {
451 debug!(?frame, "websocket user sent close frame");
452 Some(SessionExit::Loop(LoopExit::BusBreak, None))
453 }
454 Message::Text(payload) => self.handle_text(&payload).await,
455 Message::Binary(payload) => self.handle_binary(&payload).await,
456 }
457 }
458
459 async fn handle_binary(&mut self, payload: &[u8]) -> Option<SessionExit> {
460 let Ok(payload) = str::from_utf8(payload) else {
461 self.metrics.record_ws_bus_invalid_input_failure();
462 warn!("received websocket binary frame with invalid UTF-8");
463 return Some(self.client_error(CloseCode::ProtocolError).await);
464 };
465 self.handle_text(payload).await
466 }
467
468 async fn handle_text(&mut self, payload: &str) -> Option<SessionExit> {
469 let batch = match decode_client_batch(payload) {
470 Ok(batch) => batch,
471 Err(error) => {
472 let failure = error.kind();
473 match failure {
474 ClientBatchDecodeFailureKind::InvalidInput => {
475 self.metrics.record_ws_bus_invalid_input_failure();
476 }
477 ClientBatchDecodeFailureKind::UnsupportedFeature => {
478 self.metrics.record_ws_bus_unsupported_feature_failure();
479 }
480 }
481 warn!(?failure, "failed to decode client websocket batch");
482 return Some(self.client_error(CloseCode::ProtocolError).await);
483 }
484 };
485 self.metrics.record_ws_bus_batch_received(batch.len());
486 let mut output = UserOutput::new();
487 for envelope in batch {
488 if let Some(exit) = self.shutdown_exit() {
489 return Some(exit);
490 }
491 match &envelope {
492 ClientEnvelope::Request { .. } => self.metrics.record_ws_bus_client_request(),
493 ClientEnvelope::Message(_) => self.metrics.record_ws_bus_client_message(),
494 ClientEnvelope::Response { .. } => {}
495 }
496 let result = self.user.apply_client_envelope(envelope).await;
497 if let Some(exit) = self.shutdown_exit() {
498 return Some(exit);
499 }
500 match result {
501 Ok(user_output) => output.extend(user_output),
502 Err(error) => return Some(self.client_error(map_user_error(error)).await),
503 }
504 }
505 let result = send_user_output_bounded(&mut self.writer, output).await;
506 if let Some(exit) = self.shutdown_exit() {
507 return Some(exit);
508 }
509 match result {
510 Ok(_sent) => None,
511 Err(code) => Some(SessionExit::closing(LoopExit::BusBreak, code)),
512 }
513 }
514
515 async fn client_error(&self, fallback_code: CloseCode) -> SessionExit {
516 let close_code = if self.user.is_current_connection().await {
517 fallback_code
518 } else {
519 CloseCode::Kicked
520 };
521 self.shutdown_exit()
522 .unwrap_or(SessionExit::closing(LoopExit::BusBreak, close_code))
523 }
524
525 async fn handle_outbound_event(&mut self, outbound: UserOutboundEvent) -> Option<SessionExit> {
526 match outbound {
527 UserOutboundEvent::Message(UserOutbound::Close(_)) => {
528 Some(self.outbound_error(CloseCode::Kicked, false))
529 }
530 UserOutboundEvent::Message(outbound) => {
531 let result = self.user.apply_room_outbound(outbound).await;
532 if let Some(exit) = self.shutdown_exit() {
533 return Some(exit);
534 }
535 let output = match result {
536 Ok(output) => output,
537 Err(error) => {
538 return Some(self.outbound_error(map_user_error(error), false));
539 }
540 };
541 let envelope_count = output.len();
542 let result = send_user_output_bounded(&mut self.writer, output).await;
543 if let Some(exit) = self.shutdown_exit() {
544 return Some(exit);
545 }
546 match result {
547 Ok(batch_count) => {
548 self.metrics
549 .record_ws_bus_batches_sent(batch_count, envelope_count);
550 None
551 }
552 Err(code) => Some(self.outbound_error(code, true)),
553 }
554 }
555 UserOutboundEvent::Overflow(overflow) => {
556 warn!(
557 capacity = overflow.capacity(),
558 byte_capacity = overflow.byte_capacity(),
559 queued_bytes = overflow.queued_bytes(),
560 message_bytes = overflow.message_bytes(),
561 overflow_kind = ?overflow.kind(),
562 "closing websocket because the outbound queue overflowed"
563 );
564 let code = if self.user.is_current_connection().await {
566 CloseCode::Overloaded
567 } else {
568 CloseCode::Kicked
569 };
570 Some(
571 self.shutdown_exit()
572 .unwrap_or(SessionExit::closing(LoopExit::OutboundQueueOverflow, code)),
573 )
574 }
575 UserOutboundEvent::Closed => {
576 debug!("user outbound room closed");
577 Some(SessionExit::Loop(LoopExit::OutboundChannelClosed, None))
578 }
579 }
580 }
581
582 fn outbound_error(&self, code: CloseCode, log_send_failure: bool) -> SessionExit {
583 if code == CloseCode::Kicked {
584 debug!(
585 close_code = u16::from(code),
586 "closing websocket from outbound signal"
587 );
588 return SessionExit::closing(LoopExit::OutboundCloseSignal, CloseCode::Kicked);
589 }
590 self.metrics.record_ws_bus_send_failure();
591 if log_send_failure {
592 debug!(
593 close_code = u16::from(code),
594 "failed to send outbound user event"
595 );
596 }
597 SessionExit::Loop(LoopExit::OutboundMessageSendFailure, None)
598 }
599}
600
601fn map_user_error(error: UserError) -> CloseCode {
602 match error {
603 UserError::ProtocolViolation => CloseCode::ProtocolError,
604 UserError::Kicked => CloseCode::Kicked,
605 UserError::InternalError => CloseCode::Error,
606 }
607}