use std::time::{Duration, SystemTime, UNIX_EPOCH};
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine as _;
use hmac::{Hmac, Mac};
use serde_json::{json, Map, Value};
use sha2::Sha256;
use crate::error::{Error, Result};
pub const DEFAULT_TOKEN_TTL: Duration = Duration::from_secs(24 * 60 * 60);
const JWT_HEADER: &str = r#"{"alg":"HS256","typ":"JWT"}"#;
type HmacSha256 = Hmac<Sha256>;
#[derive(Debug, Clone, Default)]
pub struct GenerateTokenParams {
pub api_key: String,
pub secret_key: String,
pub permissions: Option<Vec<String>>,
pub roles: Option<Vec<String>>,
pub version: Option<i64>,
pub expires_in: Option<Duration>,
pub claims: Map<String, Value>,
}
#[derive(Debug, Clone, Default, PartialEq)]
#[non_exhaustive]
pub struct TokenClaims {
pub api_key: Option<String>,
pub permissions: Vec<String>,
pub roles: Vec<String>,
pub version: Option<i64>,
pub issued_at: Option<i64>,
pub expires_at: Option<i64>,
pub raw: Map<String, Value>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Grant {
AllowJoin,
AskJoin,
AllowMod,
}
impl Grant {
pub fn as_str(self) -> &'static str {
match self {
Grant::AllowJoin => "allow_join",
Grant::AskJoin => "ask_join",
Grant::AllowMod => "allow_mod",
}
}
}
fn unix_now() -> i64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.unwrap_or(0)
}
fn invalid_expires_in(value: &str) -> Error {
Error::validation(format!(
"invalid expiresIn {value:?}: use a number of seconds or a string like \"24h\", \"7d\""
))
}
pub fn parse_expires_in(value: &str) -> Result<Duration> {
let s = value.trim();
let digits_end = s.find(|c: char| !c.is_ascii_digit()).unwrap_or(s.len());
if digits_end == 0 {
return Err(invalid_expires_in(value));
}
let amount: u64 = s[..digits_end]
.parse()
.map_err(|_| invalid_expires_in(value))?;
let seconds_per_unit = match s[digits_end..].trim_start() {
"" | "s" => 1,
"m" => 60,
"h" => 3_600,
"d" => 86_400,
"w" => 604_800,
"y" => 31_536_000,
_ => return Err(invalid_expires_in(value)),
};
amount
.checked_mul(seconds_per_unit)
.map(Duration::from_secs)
.ok_or_else(|| invalid_expires_in(value))
}
fn new_mac(secret: &str) -> HmacSha256 {
<HmacSha256 as Mac>::new_from_slice(secret.as_bytes()).expect("HMAC accepts keys of any length")
}
fn sign_hs256(signing_input: &str, secret: &str) -> String {
let mut mac = new_mac(secret);
mac.update(signing_input.as_bytes());
URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes())
}
fn decode_segment(segment: &str) -> std::result::Result<Vec<u8>, base64::DecodeError> {
let normalized: String = segment
.trim_end_matches('=')
.chars()
.map(|c| match c {
'+' => '-',
'/' => '_',
other => other,
})
.collect();
URL_SAFE_NO_PAD.decode(normalized)
}
pub fn generate_token(params: &GenerateTokenParams) -> Result<String> {
if params.api_key.is_empty() {
return Err(Error::config("generate_token: api_key is required"));
}
if params.secret_key.is_empty() {
return Err(Error::config("generate_token: secret_key is required"));
}
let ttl = params
.expires_in
.filter(|d| !d.is_zero())
.unwrap_or(DEFAULT_TOKEN_TTL);
let iat = unix_now();
let exp = iat + ttl.as_secs() as i64;
let permissions = params
.permissions
.clone()
.unwrap_or_else(|| vec!["allow_join".to_string(), "allow_mod".to_string()]);
let roles = params
.roles
.clone()
.unwrap_or_else(|| vec!["crawler".to_string()]);
let mut payload = Map::new();
payload.insert("apikey".into(), json!(params.api_key));
payload.insert("permissions".into(), json!(permissions));
payload.insert("version".into(), json!(params.version.unwrap_or(2)));
payload.insert("roles".into(), json!(roles));
for (key, value) in ¶ms.claims {
payload.insert(key.clone(), value.clone());
}
payload.insert("iat".into(), json!(iat));
payload.insert("exp".into(), json!(exp));
let payload_json = serde_json::to_vec(&payload).map_err(|e| Error::Encode { source: e })?;
let signing_input = format!(
"{}.{}",
URL_SAFE_NO_PAD.encode(JWT_HEADER),
URL_SAFE_NO_PAD.encode(payload_json)
);
let signature = sign_hs256(&signing_input, ¶ms.secret_key);
Ok(format!("{signing_input}.{signature}"))
}
fn string_vec(value: Option<&Value>) -> Vec<String> {
value
.and_then(Value::as_array)
.map(|items| {
items
.iter()
.filter_map(Value::as_str)
.map(str::to_string)
.collect()
})
.unwrap_or_default()
}
fn claims_from_map(map: Map<String, Value>) -> TokenClaims {
TokenClaims {
api_key: map
.get("apikey")
.and_then(Value::as_str)
.map(str::to_string),
permissions: string_vec(map.get("permissions")),
roles: string_vec(map.get("roles")),
version: map.get("version").and_then(Value::as_i64),
issued_at: map.get("iat").and_then(Value::as_i64),
expires_at: map.get("exp").and_then(Value::as_i64),
raw: map,
}
}
pub fn decode_token(token: &str) -> Result<TokenClaims> {
let mut parts = token.split('.');
let (_header, payload) = (parts.next(), parts.next());
let payload = payload
.filter(|p| !p.is_empty())
.ok_or_else(|| Error::validation("malformed token: expected at least a 2-part JWT"))?;
let raw = decode_segment(payload)
.map_err(|e| Error::validation(format!("malformed token payload: {e}")))?;
let map: Map<String, Value> = serde_json::from_slice(&raw)
.map_err(|e| Error::validation(format!("malformed token payload: {e}")))?;
Ok(claims_from_map(map))
}
pub fn verify_token(token: &str, secret: &str) -> Result<TokenClaims> {
let parts: Vec<&str> = token.split('.').collect();
if parts.len() != 3 || parts.iter().any(|p| p.is_empty()) {
return Err(Error::auth("malformed token: expected a 3-part JWT"));
}
let signature = decode_segment(parts[2]).map_err(|_| Error::auth("invalid token signature"))?;
let mut mac = new_mac(secret);
mac.update(format!("{}.{}", parts[0], parts[1]).as_bytes());
mac.verify_slice(&signature)
.map_err(|_| Error::auth("invalid token signature"))?;
let claims = decode_token(token).map_err(|_| Error::auth("malformed token payload"))?;
if let Some(exp) = claims.expires_at {
if unix_now() >= exp {
return Err(Error::auth("token has expired"));
}
}
Ok(claims)
}
#[derive(Debug, Clone)]
pub struct AccessTokenBuilder {
api_key: String,
secret: String,
participant_id: Option<String>,
room_id: Option<String>,
grants: Vec<Grant>,
roles: Vec<String>,
ttl: Option<Duration>,
claims: Map<String, Value>,
}
impl AccessTokenBuilder {
pub fn new(api_key: impl Into<String>, secret: impl Into<String>) -> Self {
Self {
api_key: api_key.into(),
secret: secret.into(),
participant_id: None,
room_id: None,
grants: Vec::new(),
roles: vec!["rtc".to_string()],
ttl: None,
claims: Map::new(),
}
}
pub fn set_participant(mut self, participant_id: impl Into<String>) -> Self {
self.participant_id = Some(participant_id.into());
self
}
pub fn grant(mut self, grant: Grant) -> Self {
if !self.grants.contains(&grant) {
self.grants.push(grant);
}
self
}
pub fn grants(mut self, grants: impl IntoIterator<Item = Grant>) -> Self {
for grant in grants {
self = self.grant(grant);
}
self
}
pub fn for_room(mut self, room_id: impl Into<String>) -> Self {
self.room_id = Some(room_id.into());
self
}
pub fn for_api(mut self) -> Self {
self.roles = vec!["crawler".to_string()];
self
}
pub fn expires_in(mut self, ttl: Duration) -> Self {
self.ttl = Some(ttl);
self
}
pub fn set_claim(mut self, key: impl Into<String>, value: impl Into<Value>) -> Self {
self.claims.insert(key.into(), value.into());
self
}
pub fn to_jwt(&self) -> Result<String> {
if self.api_key.is_empty() || self.secret.is_empty() {
return Err(Error::config("access_token requires an api_key and secret"));
}
let mut claims = self.claims.clone();
if let Some(participant_id) = &self.participant_id {
claims.insert("participantId".into(), json!(participant_id));
}
if let Some(room_id) = &self.room_id {
claims.insert("roomId".into(), json!(room_id));
}
let permissions = (!self.grants.is_empty())
.then(|| self.grants.iter().map(|g| g.as_str().to_string()).collect());
generate_token(&GenerateTokenParams {
api_key: self.api_key.clone(),
secret_key: self.secret.clone(),
permissions,
roles: Some(self.roles.clone()),
version: Some(2),
expires_in: self.ttl,
claims,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn params() -> GenerateTokenParams {
GenerateTokenParams {
api_key: "test-key".into(),
secret_key: "test-secret".into(),
..Default::default()
}
}
#[test]
fn parses_duration_strings() {
assert_eq!(parse_expires_in("30").unwrap(), Duration::from_secs(30));
assert_eq!(parse_expires_in("30s").unwrap(), Duration::from_secs(30));
assert_eq!(parse_expires_in("10m").unwrap(), Duration::from_secs(600));
assert_eq!(
parse_expires_in("24h").unwrap(),
Duration::from_secs(86_400)
);
assert_eq!(
parse_expires_in("7d").unwrap(),
Duration::from_secs(604_800)
);
assert_eq!(
parse_expires_in("2w").unwrap(),
Duration::from_secs(1_209_600)
);
assert_eq!(
parse_expires_in("1y").unwrap(),
Duration::from_secs(31_536_000)
);
assert_eq!(
parse_expires_in(" 15 m ").unwrap(),
Duration::from_secs(900)
);
}
#[test]
fn rejects_bad_duration_strings() {
for bad in ["", "abc", "h", "-5s", "10x", "1.5h"] {
assert!(parse_expires_in(bad).is_err(), "expected {bad:?} to fail");
}
}
#[test]
fn generated_token_has_three_segments_and_expected_header() {
let token = generate_token(¶ms()).unwrap();
let parts: Vec<&str> = token.split('.').collect();
assert_eq!(parts.len(), 3);
let header = decode_segment(parts[0]).unwrap();
assert_eq!(String::from_utf8(header).unwrap(), JWT_HEADER);
}
#[test]
fn applies_documented_defaults() {
let claims = decode_token(&generate_token(¶ms()).unwrap()).unwrap();
assert_eq!(claims.api_key.as_deref(), Some("test-key"));
assert_eq!(claims.permissions, ["allow_join", "allow_mod"]);
assert_eq!(claims.roles, ["crawler"]);
assert_eq!(claims.version, Some(2));
let (iat, exp) = (claims.issued_at.unwrap(), claims.expires_at.unwrap());
assert_eq!(exp - iat, DEFAULT_TOKEN_TTL.as_secs() as i64);
}
#[test]
fn custom_claims_may_shadow_roles_but_never_iat_or_exp() {
let mut p = params();
p.claims.insert("roles".into(), json!(["custom"]));
p.claims.insert("exp".into(), json!(1));
p.claims.insert("iat".into(), json!(1));
p.claims.insert("participantId".into(), json!("p-1"));
let claims = decode_token(&generate_token(&p).unwrap()).unwrap();
assert_eq!(claims.roles, ["custom"], "custom claims shadow roles");
assert_eq!(
claims.raw.get("participantId").unwrap().as_str(),
Some("p-1")
);
assert_ne!(claims.expires_at, Some(1), "exp must not be overridable");
assert_ne!(claims.issued_at, Some(1), "iat must not be overridable");
}
#[test]
fn verify_round_trips_a_generated_token() {
let token = generate_token(¶ms()).unwrap();
let claims = verify_token(&token, "test-secret").unwrap();
assert_eq!(claims.api_key.as_deref(), Some("test-key"));
}
#[test]
fn verify_rejects_a_wrong_secret() {
let token = generate_token(¶ms()).unwrap();
let err = verify_token(&token, "not-the-secret").unwrap_err();
assert!(err.is_authentication());
assert!(err.to_string().contains("invalid token signature"));
}
#[test]
fn verify_rejects_a_tampered_payload() {
let token = generate_token(¶ms()).unwrap();
let mut parts: Vec<&str> = token.split('.').collect();
let forged = URL_SAFE_NO_PAD.encode(br#"{"apikey":"attacker"}"#);
parts[1] = &forged;
let tampered = parts.join(".");
assert!(verify_token(&tampered, "test-secret").is_err());
}
#[test]
fn verify_rejects_malformed_and_expired_tokens() {
assert!(verify_token("a.b", "s").unwrap_err().is_authentication());
assert!(verify_token("a..c", "s").unwrap_err().is_authentication());
let mut p = params();
p.expires_in = Some(Duration::from_secs(1));
let mut token_params = p.clone();
token_params.claims.clear();
let payload = json!({"apikey": "k", "exp": unix_now() - 3600});
let signing_input = format!(
"{}.{}",
URL_SAFE_NO_PAD.encode(JWT_HEADER),
URL_SAFE_NO_PAD.encode(serde_json::to_vec(&payload).unwrap())
);
let expired = format!("{signing_input}.{}", sign_hs256(&signing_input, "s"));
let err = verify_token(&expired, "s").unwrap_err();
assert!(err.to_string().contains("token has expired"));
}
#[test]
fn decode_does_not_verify() {
let signing_input = format!(
"{}.{}",
URL_SAFE_NO_PAD.encode(JWT_HEADER),
URL_SAFE_NO_PAD.encode(br#"{"apikey":"k"}"#)
);
let unsigned = format!("{signing_input}.garbage");
assert_eq!(
decode_token(&unsigned).unwrap().api_key.as_deref(),
Some("k")
);
assert!(verify_token(&unsigned, "s").is_err());
}
#[test]
fn decode_rejects_malformed_tokens() {
assert!(decode_token("only-one-part").is_err());
assert!(decode_token("header.").is_err());
assert!(decode_token("header.!!!not-base64!!!").is_err());
}
#[test]
fn access_token_defaults_to_a_participant_token() {
let jwt = AccessTokenBuilder::new("k", "s").to_jwt().unwrap();
let claims = decode_token(&jwt).unwrap();
assert_eq!(claims.roles, ["rtc"]);
assert_eq!(claims.permissions, ["allow_join", "allow_mod"]);
assert_eq!(claims.version, Some(2));
}
#[test]
fn for_api_switches_to_the_management_role() {
let jwt = AccessTokenBuilder::new("k", "s")
.for_api()
.to_jwt()
.unwrap();
assert_eq!(decode_token(&jwt).unwrap().roles, ["crawler"]);
}
#[test]
fn grants_are_deduplicated_and_ordered() {
let jwt = AccessTokenBuilder::new("k", "s")
.grant(Grant::AllowMod)
.grant(Grant::AllowJoin)
.grant(Grant::AllowMod)
.to_jwt()
.unwrap();
assert_eq!(
decode_token(&jwt).unwrap().permissions,
["allow_mod", "allow_join"]
);
}
#[test]
fn participant_and_room_become_claims() {
let jwt = AccessTokenBuilder::new("k", "s")
.set_participant("p-1")
.for_room("r-1")
.set_claim("custom", "v")
.to_jwt()
.unwrap();
let raw = decode_token(&jwt).unwrap().raw;
assert_eq!(raw.get("participantId").unwrap().as_str(), Some("p-1"));
assert_eq!(raw.get("roomId").unwrap().as_str(), Some("r-1"));
assert_eq!(raw.get("custom").unwrap().as_str(), Some("v"));
}
#[test]
fn generate_token_requires_credentials() {
let mut p = params();
p.api_key = String::new();
assert!(generate_token(&p)
.unwrap_err()
.to_string()
.contains("api_key"));
let mut p = params();
p.secret_key = String::new();
assert!(generate_token(&p)
.unwrap_err()
.to_string()
.contains("secret_key"));
}
#[test]
fn decode_segment_tolerates_standard_base64_alphabet() {
let raw = [0xfb_u8, 0xff, 0xbf];
let standard = base64::engine::general_purpose::STANDARD.encode(raw);
assert!(standard.contains('+') || standard.contains('/'));
assert_eq!(decode_segment(&standard).unwrap(), raw);
}
}