o_sfu/runtime/websocket_server/
handshake.rs1use std::{str, sync::Arc, time::Duration};
8
9use axum::extract::ws::{Message, WebSocket};
10use o_sfu_protocol::wire::{
11 AuthPayload, ClientEnvelope, ClientMessage, UserId, UserPermissions, WebSocketCloseCode,
12};
13use serde::Deserialize;
14use tokio::time::timeout;
15use tracing::{debug, info, warn};
16
17use super::{WsWriter, controller::WebSocketServices, io::close_writer_bounded};
18use crate::{
19 core::server::room::Room,
20 runtime::{
21 auth::{self, AuthProof, RegisteredJwtClaims, WebSocketConnectClaims},
22 telemetry::schema::event as telemetry_event,
23 websocket_server::{MAX_CLIENT_FRAME_BYTES, decode_client_batch},
24 },
25};
26
27#[derive(Deserialize)]
29struct RoomScopedConnectClaims {
30 #[serde(flatten)]
31 registered: RegisteredJwtClaims,
32 #[serde(rename = "user_id", alias = "session_id")]
33 user_id: UserId,
34 label: Option<String>,
35 permissions: Option<UserPermissions>,
36}
37
38pub(super) struct WebSocketAuth(AuthProof);
40
41pub(super) struct AuthenticatedJoin {
42 pub(super) room: Arc<Room>,
43 pub(super) claims: WebSocketConnectClaims,
44 pub(super) proof: WebSocketAuth,
45}
46
47pub(super) enum HandshakeError {
48 PeerClosed,
49 Rejected(WebSocketCloseCode),
50 Shutdown,
51}
52
53pub(super) async fn authenticate(
55 state: &WebSocketServices,
56 socket: &mut WebSocket,
57 remote_address: &str,
58) -> Result<AuthenticatedJoin, HandshakeError> {
59 let auth = receive_auth(state, socket).await;
60 if state.shutdown.is_cancelled() {
61 return Err(HandshakeError::Shutdown);
62 }
63 let auth = auth?;
64 state.metrics.record_ws_handshake_credentials_received();
65 let auth = verify_auth_payload(state, &auth, remote_address).await;
66 if state.shutdown.is_cancelled() {
67 return Err(HandshakeError::Shutdown);
68 }
69 auth.map_err(HandshakeError::Rejected)
70}
71
72async fn receive_auth(
73 state: &WebSocketServices,
74 socket: &mut WebSocket,
75) -> Result<AuthPayload, HandshakeError> {
76 tokio::select! {
77 biased;
78 () = state.shutdown.cancelled() => Err(HandshakeError::Shutdown),
79 result = timeout(
80 Duration::from_millis(state.auth.authentication_timeout_ms),
81 socket.recv(),
82 ) => match result {
83 Err(_) => {
84 debug!("timed out waiting for initial websocket auth payload");
85 Err(HandshakeError::Rejected(WebSocketCloseCode::AuthTimeout))
86 }
87 Ok(None) => Err(HandshakeError::PeerClosed),
88 Ok(Some(Err(_error))) => {
89 debug!("websocket reader returned an error before authentication completed");
90 Err(HandshakeError::Rejected(WebSocketCloseCode::Error))
91 }
92 Ok(Some(Ok(message))) => parse_auth_payload(message).map_err(HandshakeError::Rejected),
93 }
94 }
95}
96
97fn parse_auth_payload(message: Message) -> Result<AuthPayload, WebSocketCloseCode> {
98 match message {
99 Message::Text(payload) if payload.len() <= MAX_CLIENT_FRAME_BYTES => {
100 decode_auth_payload_text(&payload)
101 }
102 Message::Binary(payload) if payload.len() <= MAX_CLIENT_FRAME_BYTES => {
103 str::from_utf8(&payload)
104 .map_err(|_error| WebSocketCloseCode::ProtocolError)
105 .and_then(decode_auth_payload_text)
106 }
107 Message::Close(_) => Err(WebSocketCloseCode::Clean),
108 _ => Err(WebSocketCloseCode::ProtocolError),
109 }
110}
111
112pub fn decode_auth_payload_text(payload: &str) -> Result<AuthPayload, WebSocketCloseCode> {
118 let batch = decode_client_batch(payload).map_err(|_error| WebSocketCloseCode::ProtocolError)?;
119 let [envelope] = batch.try_into().map_err(|batch: Vec<ClientEnvelope>| {
120 warn!(
121 batch_len = batch.len(),
122 "authentication batch must contain exactly one envelope"
123 );
124 WebSocketCloseCode::ProtocolError
125 })?;
126 let ClientEnvelope::Message(ClientMessage::Auth(auth_payload)) = envelope else {
127 debug!("first websocket envelope was not an auth message");
128 return Err(WebSocketCloseCode::ProtocolError);
129 };
130 Ok(auth_payload)
131}
132
133async fn verify_auth_payload(
134 state: &WebSocketServices,
135 auth_payload: &AuthPayload,
136 remote_address: &str,
137) -> Result<AuthenticatedJoin, WebSocketCloseCode> {
138 let room = resolve_handshake_room(state, auth_payload).await?;
139 let (claims, proof) =
140 authenticate_room_scoped_claims(&auth_payload.jwt, &room, remote_address)?;
141 Ok(AuthenticatedJoin {
142 room,
143 claims,
144 proof: WebSocketAuth(proof),
145 })
146}
147
148async fn resolve_handshake_room(
150 state: &WebSocketServices,
151 auth_payload: &AuthPayload,
152) -> Result<Arc<Room>, WebSocketCloseCode> {
153 let Some(explicit_room_id) = auth_payload.channel.as_deref() else {
154 let unverified_claims = auth::decode_unverified_claims::<WebSocketConnectClaims>(
156 &auth_payload.jwt,
157 )
158 .map_err(|_error| {
159 debug!("authentication payload did not select a room");
160 WebSocketCloseCode::AuthFailed
161 })?;
162 return resolve_room_by_id(state, &unverified_claims.room_id).await;
163 };
164 resolve_room_by_id(state, explicit_room_id).await
165}
166
167async fn resolve_room_by_id(
168 state: &WebSocketServices,
169 room_id: &str,
170) -> Result<Arc<Room>, WebSocketCloseCode> {
171 state
172 .room_manager
173 .get_by_uuid(room_id)
174 .await
175 .ok_or_else(|| {
176 debug!(
177 room_id,
178 "authentication referenced an unknown explicit room"
179 );
180 WebSocketCloseCode::AuthFailed
181 })
182}
183
184fn authenticate_room_scoped_claims(
185 token: &str,
186 room: &Room,
187 remote_address: &str,
188) -> Result<(WebSocketConnectClaims, AuthProof), WebSocketCloseCode> {
189 if let Ok((mut claims, proof)) =
190 auth::verify_with_proof::<WebSocketConnectClaims>(token, room.key())
191 {
192 if claims.room_id != room.uuid() {
193 debug!(
194 expected_room_id = room.uuid(),
195 claimed_room_id = claims.room_id,
196 "room-scoped websocket token targeted the wrong room"
197 );
198 return Err(WebSocketCloseCode::AuthFailed);
199 }
200 claims.normalize_runtime_user_id();
201 return Ok((claims, proof));
202 }
203
204 let (claims, proof) = auth::verify_with_proof::<RoomScopedConnectClaims>(token, room.key())
205 .map_err(|_error| {
206 warn!(
207 remote_address,
208 "failed to verify websocket auth token against the room-scoped key"
209 );
210 WebSocketCloseCode::AuthFailed
211 })?;
212 let mut claims = WebSocketConnectClaims {
213 registered: claims.registered,
214 room_id: room.uuid().to_owned(),
215 user_id: claims.user_id,
216 label: claims.label,
217 permissions: claims.permissions,
218 };
219 claims.normalize_runtime_user_id();
220 Ok((claims, proof))
221}
222
223pub(super) async fn reject(
224 state: &WebSocketServices,
225 writer: &mut WsWriter,
226 code: WebSocketCloseCode,
227 remote_address: &str,
228 message: &str,
229) {
230 state.metrics.record_ws_handshake_rejection(Some(code));
231 info!(
232 event = telemetry_event::WS_HANDSHAKE_REJECTED,
233 close_code = u16::from(code),
234 remote_address,
235 "{message}"
236 );
237 close_writer_bounded(writer, code).await;
238}
239
240#[cfg(test)]
241#[path = "TESTS/handshake.rs"]
242mod tests;