Skip to main content

o_sfu/runtime/websocket_server/
session.rs

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    /// Checks transport health before each ping so RTC teardown closes idle sessions.
328    #[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}