Skip to main content

o_sfu/runtime/websocket_server/
session.rs

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    /// Checks transport health before each ping so RTC teardown closes idle sessions.
332    #[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                // Overflow can hide the close signal from replacement or explicit removal.
565                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}