1use std::{
2 collections::BTreeMap,
3 time::{Duration, SystemTime, UNIX_EPOCH},
4};
5
6use base64::{
7 Engine as _, decoded_len_estimate,
8 engine::general_purpose::{STANDARD, URL_SAFE},
9};
10use hmac::{Hmac, KeyInit, Mac};
11use o_sfu_protocol::wire::{UserId, UserPermissions};
12use o_sfu_rfc::jwt::{ALGORITHM_HS256, HS256_MIN_KEY_BYTES, JwtHeader, TYPE_JWT, URL_SAFE_NO_PAD};
13pub use o_sfu_rfc::jwt::{NumericDate, RegisteredJwtClaims};
14use secrecy::{ExposeSecret, SecretSlice, SecretString};
15use serde::{
16 Deserialize, Deserializer, Serialize,
17 de::{self, DeserializeOwned},
18};
19use sha2::Sha256;
20use thiserror::Error;
21use zeroize::Zeroize;
22
23type HmacSha256 = Hmac<Sha256>;
24
25pub(super) struct AuthProof(());
27
28pub const MAX_JWT_TOKEN_BYTES: usize = 16 * 1024;
29
30#[derive(Debug, Clone, PartialEq, Eq, Error)]
31pub enum AuthenticationError {
32 #[error("invalid JWT format")]
33 InvalidJwtFormat,
34 #[error("JWT token exceeds maximum byte length")]
35 TokenTooLarge { actual: usize, limit: usize },
36 #[error("invalid base64 encoding")]
37 InvalidBase64Encoding,
38 #[error("invalid JSON payload")]
39 InvalidJsonPayload,
40 #[error("token has no participant ID")]
41 MissingParticipantId,
42 #[error("token has no room ID")]
43 MissingRoomId,
44 #[error("token targets a different room")]
45 RoomMismatch,
46 #[error("unsupported JWT algorithm")]
47 UnsupportedAlgorithm,
48 #[error("invalid JWT signature")]
49 InvalidSignature,
50 #[error("token has no expiration")]
51 MissingExpiry,
52 #[error("HS256 key must contain at least {HS256_MIN_KEY_BYTES} decoded bytes")]
53 KeyTooShort,
54 #[error("token expired")]
55 TokenExpired,
56 #[error("token not valid yet")]
57 TokenNotYetValid,
58 #[error("token issued in the future")]
59 TokenIssuedInFuture,
60}
61
62const MAX_IAT_FUTURE_SKEW: Duration = Duration::from_mins(1);
68
69#[derive(Debug, Clone, Deserialize)]
70pub struct HttpRoomClaims {
71 #[serde(flatten)]
72 pub registered: RegisteredJwtClaims,
73 pub key: Option<SecretString>,
74 #[serde(rename = "keySeed")]
75 pub key_seed: Option<SecretString>,
76}
77
78#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
79pub struct HttpDisconnectClaims {
80 #[serde(flatten)]
81 pub registered: RegisteredJwtClaims,
82 #[serde(rename = "userIdsByRoom", alias = "sessionIdsByChannel")]
83 pub user_ids_by_room: BTreeMap<String, Vec<UserId>>,
84}
85
86impl HttpDisconnectClaims {
87 pub fn normalize_runtime_user_ids(&mut self) {
88 for user_ids in self.user_ids_by_room.values_mut() {
89 for user_id in user_ids {
90 *user_id = user_id.runtime_normalized();
91 }
92 }
93 }
94}
95
96#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
97pub struct WebSocketConnectClaims {
98 #[serde(flatten)]
99 pub registered: RegisteredJwtClaims,
100 pub room_id: String,
101 pub user_id: UserId,
102 #[serde(skip_serializing_if = "Option::is_none")]
103 pub label: Option<String>,
104 #[serde(skip_serializing_if = "Option::is_none")]
105 pub permissions: Option<UserPermissions>,
106}
107
108#[derive(Deserialize)]
111struct WebSocketWireClaims {
112 #[serde(flatten)]
113 registered: RegisteredJwtClaims,
114 #[serde(rename = "room_id", alias = "sfu_channel_uuid")]
115 room_id: Option<String>,
116 session_id: Option<UserId>,
117 user_id: Option<UserId>,
118 label: Option<String>,
119 permissions: Option<UserPermissions>,
120}
121
122impl WebSocketWireClaims {
123 fn into_claims(self, room_id: String) -> Result<WebSocketConnectClaims, AuthenticationError> {
124 if self
125 .room_id
126 .as_ref()
127 .is_some_and(|claimed| claimed != &room_id)
128 {
129 return Err(AuthenticationError::RoomMismatch);
130 }
131 let user_id = self
132 .session_id
133 .or(self.user_id)
134 .ok_or(AuthenticationError::MissingParticipantId)?;
135 Ok(WebSocketConnectClaims {
136 registered: self.registered,
137 room_id,
138 user_id,
139 label: self.label,
140 permissions: self.permissions,
141 })
142 }
143}
144
145impl<'de> Deserialize<'de> for WebSocketConnectClaims {
146 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
147 let mut wire = WebSocketWireClaims::deserialize(deserializer)?;
148 let room_id = wire
149 .room_id
150 .take()
151 .ok_or_else(|| de::Error::missing_field("room_id"))?;
152 wire.into_claims(room_id).map_err(de::Error::custom)
153 }
154}
155
156impl WebSocketConnectClaims {
157 pub fn normalize_runtime_user_id(&mut self) {
158 self.user_id = self.user_id.runtime_normalized();
159 }
160}
161
162#[must_use]
163pub(crate) fn duration_since_epoch() -> Duration {
164 SystemTime::now()
165 .duration_since(UNIX_EPOCH)
166 .unwrap_or(Duration::ZERO)
167}
168
169pub fn sign<T>(claims: &T, key_b64: &SecretString) -> Result<SecretString, AuthenticationError>
173where
174 T: Serialize,
175{
176 let key = decode_key(key_b64)?;
177 let header = JwtHeader {
178 alg: ALGORITHM_HS256.to_owned(),
179 typ: Some(TYPE_JWT.to_owned()),
180 };
181 let header_json =
182 serde_json::to_vec(&header).map_err(|_error| AuthenticationError::InvalidJsonPayload)?;
183 let claims_json: SecretSlice<u8> = serde_json::to_vec(claims)
186 .map_err(|_error| AuthenticationError::InvalidJsonPayload)?
187 .into();
188 let header_b64 = URL_SAFE_NO_PAD.encode(header_json);
189 let claims_b64: SecretString = URL_SAFE_NO_PAD.encode(claims_json.expose_secret()).into();
190 let signed_data: SecretString = format!("{header_b64}.{}", claims_b64.expose_secret()).into();
191 let signature = sign_hs256(signed_data.expose_secret().as_bytes(), &key)?;
192 let signature_b64 = URL_SAFE_NO_PAD.encode(signature);
193 Ok(format!("{}.{signature_b64}", signed_data.expose_secret()).into())
194}
195
196pub fn verify<T>(token: &SecretString, key_b64: &SecretString) -> Result<T, AuthenticationError>
206where
207 T: DeserializeOwned,
208{
209 validate_token_length(token.expose_secret())?;
210 verify_claims(token, &decode_key(key_b64)?)
211}
212
213pub(super) fn verify_claims<T: DeserializeOwned>(
215 token: &SecretString,
216 key: &SecretSlice<u8>,
217) -> Result<T, AuthenticationError> {
218 let token = token.expose_secret();
219 validate_token_length(token)?;
220 let (header_b64, claims_b64, signature_b64) = split_token(token)?;
221 let header_bytes = decode_jwt_segment(header_b64)?;
222 let header: JwtHeader = serde_json::from_slice(&header_bytes)
223 .map_err(|_error| AuthenticationError::InvalidJsonPayload)?;
224 if header.alg != ALGORITHM_HS256 {
225 return Err(AuthenticationError::UnsupportedAlgorithm);
226 }
227 let actual_signature = decode_jwt_segment(signature_b64)?;
228 let signed_data = &token[..token.len() - signature_b64.len() - 1];
231 verify_hs256(signed_data.as_bytes(), key, &actual_signature)?;
232 let claims_bytes: SecretSlice<u8> = decode_jwt_segment(claims_b64)?.into();
233 let registered_claims: RegisteredJwtClaims =
234 serde_json::from_slice(claims_bytes.expose_secret())
235 .map_err(|_error| AuthenticationError::InvalidJsonPayload)?;
236 validate_registered_claims(®istered_claims)?;
237 serde_json::from_slice(claims_bytes.expose_secret())
238 .map_err(|_error| AuthenticationError::InvalidJsonPayload)
239}
240
241pub(super) fn verify_websocket_claims(
246 token: &SecretString,
247 key: &SecretSlice<u8>,
248 room_id: &str,
249) -> Result<(WebSocketConnectClaims, AuthProof), AuthenticationError> {
250 let wire: WebSocketWireClaims = verify_claims(token, key)?;
251 let mut claims = wire.into_claims(room_id.to_owned())?;
252 claims.normalize_runtime_user_id();
253 Ok((claims, AuthProof(())))
254}
255
256pub(super) fn unverified_websocket_room(
261 token: &SecretString,
262) -> Result<String, AuthenticationError> {
263 let token = token.expose_secret();
264 validate_token_length(token)?;
265 let (_header_b64, claims_b64, _signature_b64) = split_token(token)?;
266 let claims_bytes: SecretSlice<u8> = decode_jwt_segment(claims_b64)?.into();
267 let claims: WebSocketWireClaims = serde_json::from_slice(claims_bytes.expose_secret())
268 .map_err(|_error| AuthenticationError::InvalidJsonPayload)?;
269 claims.room_id.ok_or(AuthenticationError::MissingRoomId)
270}
271
272fn validate_token_length(token: &str) -> Result<(), AuthenticationError> {
273 if token.len() > MAX_JWT_TOKEN_BYTES {
274 return Err(AuthenticationError::TokenTooLarge {
275 actual: token.len(),
276 limit: MAX_JWT_TOKEN_BYTES,
277 });
278 }
279 Ok(())
280}
281
282fn validate_registered_claims(claims: &RegisteredJwtClaims) -> Result<(), AuthenticationError> {
283 validate_registered_claims_at(claims, duration_since_epoch())
284}
285
286fn validate_registered_claims_at(
287 claims: &RegisteredJwtClaims,
288 now: Duration,
289) -> Result<(), AuthenticationError> {
290 let iat_limit = NumericDate::from(now.saturating_add(MAX_IAT_FUTURE_SKEW));
291 let now = NumericDate::from(now);
292 let exp = claims.exp.ok_or(AuthenticationError::MissingExpiry)?;
295 if exp <= now {
296 return Err(AuthenticationError::TokenExpired);
297 }
298 if claims.nbf.is_some_and(|nbf| nbf > now) {
299 return Err(AuthenticationError::TokenNotYetValid);
300 }
301 if claims.iat.is_some_and(|iat| iat > iat_limit) {
302 return Err(AuthenticationError::TokenIssuedInFuture);
303 }
304 Ok(())
305}
306
307fn sign_hs256(data: &[u8], key: &SecretSlice<u8>) -> Result<[u8; 32], AuthenticationError> {
308 let mut mac = {
309 let raw_key = key.expose_secret();
310 HmacSha256::new_from_slice(raw_key)
311 .map_err(|_error| AuthenticationError::InvalidBase64Encoding)?
312 };
313 mac.update(data);
314 Ok(mac.finalize().into_bytes().into())
315}
316
317fn verify_hs256(
318 data: &[u8],
319 key: &SecretSlice<u8>,
320 signature: &[u8],
321) -> Result<(), AuthenticationError> {
322 let mut mac = {
323 let raw_key = key.expose_secret();
324 HmacSha256::new_from_slice(raw_key)
325 .map_err(|_error| AuthenticationError::InvalidBase64Encoding)?
326 };
327 mac.update(data);
328 mac.verify_slice(signature)
329 .map_err(|_error| AuthenticationError::InvalidSignature)
330}
331
332fn split_token(token: &str) -> Result<(&str, &str, &str), AuthenticationError> {
333 let mut parts = token.split('.');
334 let (Some(header), Some(claims), Some(signature), None) =
335 (parts.next(), parts.next(), parts.next(), parts.next())
336 else {
337 return Err(AuthenticationError::InvalidJwtFormat);
338 };
339 if header.is_empty() || claims.is_empty() || signature.is_empty() {
340 return Err(AuthenticationError::InvalidJwtFormat);
341 }
342 Ok((header, claims, signature))
343}
344
345pub(crate) fn decode_key(input: &SecretString) -> Result<SecretSlice<u8>, AuthenticationError> {
346 let padded: SecretString = pad_base64(input.expose_secret()).into();
347 let padded_bytes = padded.expose_secret().as_bytes();
348 let mut buffer = vec![0u8; decoded_len_estimate(padded_bytes.len())];
351 let bytes_written = URL_SAFE
352 .decode_slice(padded_bytes, &mut buffer)
353 .or_else(|_error| STANDARD.decode_slice(padded_bytes, &mut buffer));
354 let bytes_written = match bytes_written {
355 Ok(bytes_written) => bytes_written,
356 Err(_error) => {
357 buffer.zeroize();
358 return Err(AuthenticationError::InvalidBase64Encoding);
359 }
360 };
361 buffer.truncate(bytes_written);
362 Ok(SecretSlice::from(buffer))
363}
364
365pub(crate) fn decode_signing_key(
370 input: &SecretString,
371) -> Result<SecretSlice<u8>, AuthenticationError> {
372 let key = decode_key(input)?;
373 if key.expose_secret().len() < HS256_MIN_KEY_BYTES {
374 return Err(AuthenticationError::KeyTooShort);
375 }
376 Ok(key)
377}
378
379fn decode_jwt_segment(input: &str) -> Result<Vec<u8>, AuthenticationError> {
380 URL_SAFE_NO_PAD
381 .decode(input.as_bytes())
382 .map_err(|_error| AuthenticationError::InvalidBase64Encoding)
383}
384
385fn pad_base64(input: &str) -> String {
386 let remainder = input.len() % 4;
387 if remainder == 0 {
388 return input.to_owned();
389 }
390 format!("{input}{}", "=".repeat(4 - remainder))
391}
392
393pub(crate) fn derive_key_from_seed(
394 key: &SecretSlice<u8>,
395 seed: &SecretString,
396) -> Result<SecretSlice<u8>, AuthenticationError> {
397 let seed_bytes = decode_key(seed)?;
398 let mut derived_key = sign_hs256(seed_bytes.expose_secret(), key)?;
399 let secret = SecretSlice::from(derived_key.to_vec());
400 derived_key.zeroize();
401 Ok(secret)
402}
403
404#[cfg(any(test, feature = "testing"))]
414pub mod test_support {
415 use secrecy::ExposeSecret;
416 use serde::Serialize;
417
418 use super::{HttpRoomClaims, RegisteredJwtClaims};
419
420 #[derive(Debug, Clone, PartialEq, Eq, Serialize)]
421 pub struct TestHttpRoomClaims<'a> {
422 #[serde(flatten)]
423 pub registered: RegisteredJwtClaims,
424 #[serde(skip_serializing_if = "Option::is_none")]
425 pub key: Option<&'a str>,
426 #[serde(rename = "keySeed", skip_serializing_if = "Option::is_none")]
427 pub key_seed: Option<&'a str>,
428 }
429
430 impl<'a> From<&'a HttpRoomClaims> for TestHttpRoomClaims<'a> {
431 fn from(claims: &'a HttpRoomClaims) -> Self {
432 Self {
433 registered: claims.registered.clone(),
434 key: claims.key.as_ref().map(ExposeSecret::expose_secret),
435 key_seed: claims.key_seed.as_ref().map(ExposeSecret::expose_secret),
436 }
437 }
438 }
439
440 impl PartialEq<HttpRoomClaims> for TestHttpRoomClaims<'_> {
441 fn eq(&self, other: &HttpRoomClaims) -> bool {
442 self.registered == other.registered
443 && self.key == other.key.as_ref().map(ExposeSecret::expose_secret)
444 && self.key_seed == other.key_seed.as_ref().map(ExposeSecret::expose_secret)
445 }
446 }
447
448 impl PartialEq<TestHttpRoomClaims<'_>> for HttpRoomClaims {
449 fn eq(&self, other: &TestHttpRoomClaims<'_>) -> bool {
450 other == self
451 }
452 }
453}
454
455#[cfg(test)]
456#[path = "TESTS/auth.rs"]
457mod tests;