o_sfu/runtime/websocket_server/
handshake.rs1use std::{borrow::Cow, fmt::Display, str, sync::Arc};
8
9use axum::extract::ws::{Message, WebSocket};
10use o_sfu_protocol::wire::{AuthPayload, ClientEnvelope, ClientMessage, WebSocketCloseCode};
11use thiserror::Error;
12use tokio::time::timeout;
13use tracing::info;
14
15use super::{
16 WsWriter, admission::admit_rejection_log, controller::WebSocketServices,
17 io::close_writer_bounded,
18};
19use crate::{
20 core::server::room::Room,
21 runtime::{
22 auth::{self, AuthProof, AuthenticationError, WebSocketConnectClaims},
23 telemetry::schema::event as telemetry_event,
24 websocket_server::{MAX_CLIENT_FRAME_BYTES, decode_client_batch},
25 },
26};
27
28pub(super) struct WebSocketAuth(AuthProof);
30
31pub(super) struct AuthenticatedJoin {
32 pub(super) room: Arc<Room>,
33 pub(super) claims: WebSocketConnectClaims,
34 pub(super) proof: WebSocketAuth,
35}
36
37#[derive(Debug, Error)]
38pub(super) enum HandshakeError {
39 #[error("peer closed before authentication")]
40 PeerClosed,
41 #[error("websocket handshake rejected: {0:?}")]
42 Rejected(WebSocketCloseCode),
43 #[error(transparent)]
44 Authentication(#[from] AuthenticationError),
45 #[error("authentication referenced an unknown room")]
46 UnknownRoom,
47 #[error("server is shutting down")]
48 Shutdown,
49}
50
51impl HandshakeError {
52 pub(super) fn close_code(&self) -> WebSocketCloseCode {
53 match self {
54 Self::PeerClosed => WebSocketCloseCode::Clean,
55 Self::Rejected(code) => *code,
56 Self::Authentication(_) | Self::UnknownRoom => WebSocketCloseCode::AuthFailed,
57 Self::Shutdown => WebSocketCloseCode::Leaving,
58 }
59 }
60}
61
62pub(super) async fn authenticate(
64 state: &WebSocketServices,
65 socket: &mut WebSocket,
66) -> Result<AuthenticatedJoin, HandshakeError> {
67 let auth = receive_auth(state, socket).await;
68 if state.shutdown.is_cancelled() {
69 return Err(HandshakeError::Shutdown);
70 }
71 let auth = auth?;
72 state.metrics.record_ws_handshake_credentials_received();
73 let auth = verify_auth_payload(state, &auth).await;
74 if state.shutdown.is_cancelled() {
75 return Err(HandshakeError::Shutdown);
76 }
77 auth
78}
79
80async fn receive_auth(
81 state: &WebSocketServices,
82 socket: &mut WebSocket,
83) -> Result<AuthPayload, HandshakeError> {
84 tokio::select! {
85 biased;
86 () = state.shutdown.cancelled() => Err(HandshakeError::Shutdown),
87 result = timeout(
88 state.authentication_timeout.as_duration(),
89 socket.recv(),
90 ) => match result {
91 Err(_) => Err(HandshakeError::Rejected(WebSocketCloseCode::AuthTimeout)),
92 Ok(None) => Err(HandshakeError::PeerClosed),
93 Ok(Some(Err(_error))) => {
94 Err(HandshakeError::Rejected(WebSocketCloseCode::Error))
95 }
96 Ok(Some(Ok(message))) => parse_auth_payload(message).map_err(HandshakeError::Rejected),
97 }
98 }
99}
100
101fn parse_auth_payload(message: Message) -> Result<AuthPayload, WebSocketCloseCode> {
102 match message {
103 Message::Text(payload) if payload.len() <= MAX_CLIENT_FRAME_BYTES => {
104 decode_auth_payload_text(&payload)
105 }
106 Message::Binary(payload) if payload.len() <= MAX_CLIENT_FRAME_BYTES => {
107 str::from_utf8(&payload)
108 .map_err(|_error| WebSocketCloseCode::ProtocolError)
109 .and_then(decode_auth_payload_text)
110 }
111 Message::Close(_) => Err(WebSocketCloseCode::Clean),
112 _ => Err(WebSocketCloseCode::ProtocolError),
113 }
114}
115
116pub fn decode_auth_payload_text(payload: &str) -> Result<AuthPayload, WebSocketCloseCode> {
122 let batch = decode_client_batch(payload).map_err(|_error| WebSocketCloseCode::ProtocolError)?;
123 let [envelope] = batch
124 .try_into()
125 .map_err(|_batch: Vec<ClientEnvelope>| WebSocketCloseCode::ProtocolError)?;
126 let ClientEnvelope::Message(ClientMessage::Auth(auth_payload)) = envelope else {
127 return Err(WebSocketCloseCode::ProtocolError);
128 };
129 Ok(auth_payload)
130}
131
132async fn verify_auth_payload(
133 state: &WebSocketServices,
134 auth_payload: &AuthPayload,
135) -> Result<AuthenticatedJoin, HandshakeError> {
136 let room_id = match &auth_payload.channel {
137 Some(room_id) => Cow::Borrowed(room_id.as_str()),
138 None => Cow::Owned(auth::unverified_websocket_room(&auth_payload.jwt)?),
139 };
140 let room = state
141 .room_manager
142 .get_by_uuid(&room_id)
143 .await
144 .ok_or(HandshakeError::UnknownRoom)?;
145 let (claims, proof) =
146 auth::verify_websocket_claims(&auth_payload.jwt, room.key(), room.uuid())?;
147 Ok(AuthenticatedJoin {
148 room,
149 claims,
150 proof: WebSocketAuth(proof),
151 })
152}
153
154pub(super) async fn reject(
155 state: &WebSocketServices,
156 writer: &mut WsWriter,
157 code: WebSocketCloseCode,
158 remote_address: &str,
159 reason: impl Display,
160) {
161 state.metrics.record_ws_handshake_rejection(Some(code));
162 if let Some(suppressed_rejections) = admit_rejection_log() {
163 info!(
164 event = telemetry_event::WS_HANDSHAKE_REJECTED,
165 close_code = u16::from(code),
166 remote_address,
167 suppressed_rejections,
168 reason = %reason,
169 "rejecting websocket handshake"
170 );
171 }
172 close_writer_bounded(writer, code).await;
173}
174
175#[cfg(test)]
176#[path = "TESTS/handshake.rs"]
177mod tests;