use crate::crypto::{HmacSha256, HmacSha512};
use crate::encoding::base64url_encode;
use crate::json::escape_json_string;
use crate::util::log::debug;
use super::header::JwtAlgorithm;
#[doc(alias = "signing_algorithm")]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum JwtSigningAlgorithm {
Hs256,
Hs512,
}
impl JwtSigningAlgorithm {
#[must_use]
#[inline]
pub fn as_str(self) -> &'static str {
match self {
Self::Hs256 => "HS256",
Self::Hs512 => "HS512",
}
}
#[must_use]
#[inline]
pub fn to_jwt_algorithm(self) -> JwtAlgorithm {
match self {
Self::Hs256 => JwtAlgorithm::HS256,
Self::Hs512 => JwtAlgorithm::HS512,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum CustomClaimValue {
Str(String),
Bool(bool),
Number(u64),
StrArray(Vec<String>),
}
pub trait IntoCustomClaim {
#[doc(hidden)]
fn into_custom_claim(self) -> CustomClaimValueOpaque;
}
#[doc(hidden)]
#[derive(Debug, Clone)]
pub struct CustomClaimValueOpaque(CustomClaimValue);
impl IntoCustomClaim for &str {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::Str(self.to_owned()))
}
}
impl IntoCustomClaim for CustomClaimValueOpaque {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
self
}
}
impl IntoCustomClaim for String {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::Str(self))
}
}
impl IntoCustomClaim for bool {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::Bool(self))
}
}
impl IntoCustomClaim for u64 {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::Number(self))
}
}
impl IntoCustomClaim for Vec<String> {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::StrArray(self))
}
}
impl IntoCustomClaim for Vec<&str> {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::StrArray(
self.into_iter().map(ToOwned::to_owned).collect(),
))
}
}
impl IntoCustomClaim for &[&str] {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::StrArray(
self.iter().map(|s| (*s).to_owned()).collect(),
))
}
}
impl<const N: usize> IntoCustomClaim for [&str; N] {
fn into_custom_claim(self) -> CustomClaimValueOpaque {
CustomClaimValueOpaque(CustomClaimValue::StrArray(
self.iter().map(|s| (*s).to_owned()).collect(),
))
}
}
#[doc(alias = "jwt_builder")]
#[derive(Debug, Clone, Default)]
#[must_use = "an encoder does nothing until `sign` is called"]
pub struct JwtEncoder {
iss: Option<String>,
sub: Option<String>,
aud: Vec<String>,
exp: Option<u64>,
nbf: Option<u64>,
iat: Option<u64>,
jti: Option<String>,
nonce: Option<String>,
kid: Option<String>,
custom: Vec<(String, CustomClaimValue)>,
}
impl JwtEncoder {
#[inline]
pub fn new() -> Self {
Self::default()
}
#[inline]
pub fn issuer(mut self, iss: impl Into<String>) -> Self {
self.iss = Some(iss.into());
self
}
#[inline]
pub fn subject(mut self, sub: impl Into<String>) -> Self {
self.sub = Some(sub.into());
self
}
#[inline]
pub fn audience(mut self, aud: impl Into<String>) -> Self {
self.aud.push(aud.into());
self
}
#[inline]
pub fn expiration(mut self, exp: u64) -> Self {
self.exp = Some(exp);
self
}
#[inline]
pub fn not_before(mut self, nbf: u64) -> Self {
self.nbf = Some(nbf);
self
}
#[inline]
pub fn issued_at(mut self, iat: u64) -> Self {
self.iat = Some(iat);
self
}
#[inline]
pub fn jwt_id(mut self, jti: impl Into<String>) -> Self {
self.jti = Some(jti.into());
self
}
#[inline]
pub fn nonce(mut self, nonce: impl Into<String>) -> Self {
self.nonce = Some(nonce.into());
self
}
#[inline]
pub fn key_id(mut self, kid: impl Into<String>) -> Self {
self.kid = Some(kid.into());
self
}
#[inline]
pub fn custom_claim(mut self, name: impl Into<String>, value: impl IntoCustomClaim) -> Self {
self.custom.push((name.into(), value.into_custom_claim().0));
self
}
#[must_use = "this returns the signed JWT string"]
pub fn sign(&self, alg: JwtSigningAlgorithm, key: &[u8]) -> String {
let signing_input = self.signing_input(alg.as_str(), self.kid.as_deref());
let signature_b64 = match alg {
JwtSigningAlgorithm::Hs256 => {
base64url_encode(&HmacSha256::mac(key, signing_input.as_bytes()))
}
JwtSigningAlgorithm::Hs512 => {
base64url_encode(&HmacSha512::mac(key, signing_input.as_bytes()))
}
};
debug!(alg = %alg.as_str(), "jwt: token minted");
let mut token = signing_input;
token.push('.');
token.push_str(&signature_b64);
token
}
#[cfg(feature = "asym-jwt")]
#[must_use = "this returns the signed JWT string"]
pub fn sign_asymmetric(&self, key: &crate::jwt::AsymmetricSigningKey) -> String {
let kid = self.kid.as_deref().unwrap_or_else(|| key.kid());
let alg = key.algorithm().as_str();
let signing_input = self.signing_input(alg, Some(kid));
let signature_b64 = base64url_encode(&key.sign(signing_input.as_bytes()));
debug!(alg = %alg, "jwt: token minted (asymmetric)");
let mut token = signing_input;
token.push('.');
token.push_str(&signature_b64);
token
}
fn signing_input(&self, alg: &str, kid: Option<&str>) -> String {
let header_json = match kid {
Some(kid) => format!(
r#"{{"alg":"{alg}","typ":"JWT","kid":{}}}"#,
escape_json_string(kid),
),
None => format!(r#"{{"alg":"{alg}","typ":"JWT"}}"#),
};
let header_b64 = base64url_encode(header_json.as_bytes());
let payload_json = self.build_payload_json();
let payload_b64 = base64url_encode(payload_json.as_bytes());
super::signature::signing_input(&header_b64, &payload_b64)
}
fn build_payload_json(&self) -> String {
let mut parts: Vec<String> = Vec::new();
let mut seen: Vec<&str> = Vec::new();
if let Some(iss) = &self.iss {
parts.push(format!(r#""iss":{}"#, escape_json_string(iss)));
seen.push("iss");
}
if let Some(sub) = &self.sub {
parts.push(format!(r#""sub":{}"#, escape_json_string(sub)));
seen.push("sub");
}
match self.aud.as_slice() {
[] => {}
[single] => {
parts.push(format!(r#""aud":{}"#, escape_json_string(single)));
seen.push("aud");
}
multiple => {
let items: Vec<String> = multiple.iter().map(|a| escape_json_string(a)).collect();
parts.push(format!(r#""aud":[{}]"#, items.join(",")));
seen.push("aud");
}
}
if let Some(exp) = self.exp {
parts.push(format!(r#""exp":{exp}"#));
seen.push("exp");
}
if let Some(nbf) = self.nbf {
parts.push(format!(r#""nbf":{nbf}"#));
seen.push("nbf");
}
if let Some(iat) = self.iat {
parts.push(format!(r#""iat":{iat}"#));
seen.push("iat");
}
if let Some(jti) = &self.jti {
parts.push(format!(r#""jti":{}"#, escape_json_string(jti)));
seen.push("jti");
}
if let Some(nonce) = &self.nonce {
parts.push(format!(r#""nonce":{}"#, escape_json_string(nonce)));
seen.push("nonce");
}
for (name, value) in &self.custom {
if seen.contains(&name.as_str()) {
continue;
}
seen.push(name);
let rendered = match value {
CustomClaimValue::Str(s) => escape_json_string(s),
CustomClaimValue::Bool(b) => b.to_string(),
CustomClaimValue::Number(n) => n.to_string(),
CustomClaimValue::StrArray(items) => {
let rendered: Vec<String> =
items.iter().map(|s| escape_json_string(s)).collect();
format!("[{}]", rendered.join(","))
}
};
parts.push(format!("{}:{rendered}", escape_json_string(name)));
}
format!("{{{}}}", parts.join(","))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::jwt::{JwtAlgorithm, verify_jwt};
const KEY: &[u8] = b"super-secret-key-for-testing-only";
#[test]
fn minted_token_has_three_segments() {
let token = JwtEncoder::new().sign(JwtSigningAlgorithm::Hs256, KEY);
assert_eq!(token.split('.').count(), 3);
}
#[test]
fn empty_encoder_produces_empty_payload_object() {
let token = JwtEncoder::new().sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(claims.iss(), None);
assert_eq!(claims.sub(), None);
assert!(claims.aud().is_empty());
}
#[test]
fn custom_claim_colliding_with_registered_is_dropped() {
let token = JwtEncoder::new()
.subject("real-subject")
.custom_claim("sub", "attacker")
.custom_claim("role", "admin")
.custom_claim("role", "user")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(claims.sub(), Some("real-subject"));
assert_eq!(
claims.get_claim("role").and_then(|v| v.as_str()),
Some("admin"),
"first custom occurrence wins",
);
}
#[test]
fn round_trip_hs256() {
let token = JwtEncoder::new()
.issuer("https://auth.example.com")
.subject("user-123")
.audience("my-client")
.expiration(9_999_999_999)
.not_before(1_000)
.issued_at(1_700_000_000)
.jwt_id("token-abc")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (header, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(header.alg(), JwtAlgorithm::HS256);
assert_eq!(header.typ(), Some("JWT"));
assert_eq!(claims.iss(), Some("https://auth.example.com"));
assert_eq!(claims.sub(), Some("user-123"));
assert_eq!(claims.aud(), &["my-client"]);
assert_eq!(claims.exp(), Some(9_999_999_999));
assert_eq!(claims.nbf(), Some(1_000));
assert_eq!(claims.iat(), Some(1_700_000_000));
assert_eq!(claims.jti(), Some("token-abc"));
}
#[test]
fn round_trip_hs512() {
let token = JwtEncoder::new()
.subject("user-9")
.sign(JwtSigningAlgorithm::Hs512, KEY);
let (header, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(header.alg(), JwtAlgorithm::HS512);
assert_eq!(claims.sub(), Some("user-9"));
}
#[test]
fn wrong_key_fails_verification() {
let token = JwtEncoder::new()
.subject("user-1")
.sign(JwtSigningAlgorithm::Hs256, KEY);
assert!(verify_jwt(&token, b"wrong-key").is_err());
}
#[test]
fn single_audience_is_emitted_as_string() {
let token = JwtEncoder::new()
.audience("only-client")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let payload = token.split('.').nth(1).unwrap();
let bytes = crate::encoding::base64url_decode(payload).unwrap();
let json = std::str::from_utf8(&bytes).unwrap();
assert!(json.contains(r#""aud":"only-client""#), "got: {json}");
}
#[test]
fn multiple_audiences_are_emitted_as_array() {
let token = JwtEncoder::new()
.audience("client-a")
.audience("client-b")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(claims.aud(), &["client-a", "client-b"]);
let payload = token.split('.').nth(1).unwrap();
let bytes = crate::encoding::base64url_decode(payload).unwrap();
let json = std::str::from_utf8(&bytes).unwrap();
assert!(
json.contains(r#""aud":["client-a","client-b"]"#),
"got: {json}"
);
}
#[test]
fn custom_string_claim() {
let token = JwtEncoder::new()
.custom_claim("preferred_username", "frodo")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(
claims
.get_claim("preferred_username")
.and_then(|v| v.as_str()),
Some("frodo"),
);
}
#[test]
fn custom_bool_claim() {
let token = JwtEncoder::new()
.custom_claim("email_verified", true)
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(
claims
.get_claim("email_verified")
.and_then(crate::json::JsonValue::as_bool),
Some(true),
);
}
#[test]
fn custom_string_array_claim() {
let token = JwtEncoder::new()
.custom_claim("roles", ["admin", "developer"])
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
let roles = claims
.get_claim("roles")
.and_then(|v| v.as_array())
.unwrap();
let names: Vec<&str> = roles.iter().filter_map(|v| v.as_str()).collect();
assert_eq!(names, ["admin", "developer"]);
}
#[test]
fn custom_claim_accepts_owned_string_and_vec() {
let token = JwtEncoder::new()
.custom_claim("tenant".to_string(), String::from("entropy"))
.custom_claim("groups", vec!["a".to_string(), "b".to_string()])
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(
claims.get_claim("tenant").and_then(|v| v.as_str()),
Some("entropy")
);
assert_eq!(
claims
.get_claim("groups")
.and_then(|v| v.as_array())
.unwrap()
.len(),
2
);
}
#[test]
fn nonce_claim_round_trips() {
let token = JwtEncoder::new()
.nonce("n-0S6_WzA2Mj")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(
claims.get_claim("nonce").and_then(|v| v.as_str()),
Some("n-0S6_WzA2Mj")
);
}
#[test]
fn escapes_special_characters_in_string_claims() {
let token = JwtEncoder::new()
.subject("a\"b\\c\nd")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(claims.sub(), Some("a\"b\\c\nd"));
}
#[test]
fn escapes_control_characters() {
let token = JwtEncoder::new()
.subject("tab\there")
.custom_claim("x", "\u{1}\u{1f}")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(claims.sub(), Some("tab\there"));
assert_eq!(
claims.get_claim("x").and_then(|v| v.as_str()),
Some("\u{1}\u{1f}")
);
}
#[test]
fn preserves_non_ascii_unicode() {
let token = JwtEncoder::new()
.custom_claim("name", "Sméagol 🧙")
.sign(JwtSigningAlgorithm::Hs256, KEY);
let (_, claims) = verify_jwt(&token, KEY).unwrap();
assert_eq!(
claims.get_claim("name").and_then(|v| v.as_str()),
Some("Sméagol 🧙")
);
}
#[test]
fn signing_algorithm_as_str() {
assert_eq!(JwtSigningAlgorithm::Hs256.as_str(), "HS256");
assert_eq!(JwtSigningAlgorithm::Hs512.as_str(), "HS512");
}
#[test]
fn signing_algorithm_maps_to_jwt_algorithm() {
assert_eq!(
JwtSigningAlgorithm::Hs256.to_jwt_algorithm(),
JwtAlgorithm::HS256
);
assert_eq!(
JwtSigningAlgorithm::Hs512.to_jwt_algorithm(),
JwtAlgorithm::HS512
);
}
#[test]
fn header_segment_matches_expected_for_hs256() {
let token = JwtEncoder::new().sign(JwtSigningAlgorithm::Hs256, KEY);
let header_b64 = token.split('.').next().unwrap();
let decoded = crate::encoding::base64url_decode(header_b64).unwrap();
assert_eq!(
std::str::from_utf8(&decoded).unwrap(),
r#"{"alg":"HS256","typ":"JWT"}"#,
);
}
}