#![deny(missing_debug_implementations, missing_docs, bare_trait_objects)]
use serde::{de::DeserializeOwned, Deserialize, Serialize};
use smallvec::{smallvec, SmallVec};
use std::{borrow::Cow, convert::TryFrom, fmt};
pub mod alg;
mod claims;
mod error;
pub use crate::{
claims::{Claims, Empty, TimeOptions},
error::{CreationError, ParseError, ValidationError},
};
pub mod prelude {
pub use crate::{AlgorithmExt as _, Claims, Header, TimeOptions, Token, UntrustedToken};
}
const SIGNATURE_SIZE: usize = 128;
pub trait AlgorithmSignature: Sized {
fn try_from_slice(slice: &[u8]) -> anyhow::Result<Self>;
fn as_bytes(&self) -> Cow<'_, [u8]>;
}
pub trait Algorithm {
type SigningKey;
type VerifyingKey;
type Signature: AlgorithmSignature;
fn name(&self) -> Cow<'static, str>;
fn sign(&self, signing_key: &Self::SigningKey, message: &[u8]) -> Self::Signature;
fn verify_signature(
&self,
signature: &Self::Signature,
verifying_key: &Self::VerifyingKey,
message: &[u8],
) -> bool;
}
#[derive(Debug, Clone, Copy)]
pub struct Renamed<A> {
inner: A,
name: &'static str,
}
impl<A: Algorithm> Renamed<A> {
pub fn new(algorithm: A, new_name: &'static str) -> Self {
Self {
inner: algorithm,
name: new_name,
}
}
}
impl<A: Algorithm> Algorithm for Renamed<A> {
type SigningKey = A::SigningKey;
type VerifyingKey = A::VerifyingKey;
type Signature = A::Signature;
fn name(&self) -> Cow<'static, str> {
Cow::Borrowed(self.name)
}
fn sign(&self, signing_key: &Self::SigningKey, message: &[u8]) -> Self::Signature {
self.inner.sign(signing_key, message)
}
fn verify_signature(
&self,
signature: &Self::Signature,
verifying_key: &Self::VerifyingKey,
message: &[u8],
) -> bool {
self.inner
.verify_signature(signature, verifying_key, message)
}
}
pub trait AlgorithmExt: Algorithm {
fn token<T>(
&self,
header: Header,
claims: &Claims<T>,
signing_key: &Self::SigningKey,
) -> Result<String, CreationError>
where
T: Serialize;
fn compact_token<T>(
&self,
header: Header,
claims: &Claims<T>,
signing_key: &Self::SigningKey,
) -> Result<String, CreationError>
where
T: Serialize;
fn validate_integrity<T>(
&self,
token: &UntrustedToken<'_>,
verifying_key: &Self::VerifyingKey,
) -> Result<Token<T>, ValidationError>
where
T: DeserializeOwned;
fn validate_for_signed_token<T>(
&self,
token: &UntrustedToken<'_>,
verifying_key: &Self::VerifyingKey,
) -> Result<SignedToken<Self, T>, ValidationError>
where
T: DeserializeOwned;
}
impl<A: Algorithm> AlgorithmExt for A {
fn token<T>(
&self,
header: Header,
claims: &Claims<T>,
signing_key: &Self::SigningKey,
) -> Result<String, CreationError>
where
T: Serialize,
{
let complete_header = CompleteHeader {
algorithm: self.name(),
content_type: None,
inner: header,
};
let header = serde_json::to_string(&complete_header).map_err(CreationError::Header)?;
let mut buffer = base64::encode_config(&header, base64::URL_SAFE_NO_PAD);
buffer.push('.');
let claims = serde_json::to_string(claims).map_err(CreationError::Claims)?;
base64::encode_config_buf(&claims, base64::URL_SAFE_NO_PAD, &mut buffer);
let signature = self.sign(signing_key, buffer.as_bytes());
buffer.push('.');
base64::encode_config_buf(
signature.as_bytes().as_ref(),
base64::URL_SAFE_NO_PAD,
&mut buffer,
);
Ok(buffer)
}
fn compact_token<T>(
&self,
header: Header,
claims: &Claims<T>,
signing_key: &Self::SigningKey,
) -> Result<String, CreationError>
where
T: Serialize,
{
let complete_header = CompleteHeader {
algorithm: self.name(),
content_type: Some("CBOR".to_owned()),
inner: header,
};
let header = serde_json::to_string(&complete_header).map_err(CreationError::Header)?;
let mut buffer = base64::encode_config(&header, base64::URL_SAFE_NO_PAD);
buffer.push('.');
let claims = serde_cbor::to_vec(claims).map_err(CreationError::CborClaims)?;
base64::encode_config_buf(&claims, base64::URL_SAFE_NO_PAD, &mut buffer);
let signature = self.sign(signing_key, buffer.as_bytes());
buffer.push('.');
base64::encode_config_buf(
signature.as_bytes().as_ref(),
base64::URL_SAFE_NO_PAD,
&mut buffer,
);
Ok(buffer)
}
fn validate_integrity<T>(
&self,
token: &UntrustedToken<'_>,
verifying_key: &Self::VerifyingKey,
) -> Result<Token<T>, ValidationError>
where
T: DeserializeOwned,
{
self.validate_for_signed_token(token, verifying_key)
.map(|wrapper| wrapper.token)
}
fn validate_for_signed_token<T>(
&self,
token: &UntrustedToken<'_>,
verifying_key: &Self::VerifyingKey,
) -> Result<SignedToken<Self, T>, ValidationError>
where
T: DeserializeOwned,
{
if self.name() != token.algorithm {
return Err(ValidationError::AlgorithmMismatch);
}
let signature = Self::Signature::try_from_slice(&token.signature[..])
.map_err(ValidationError::MalformedSignature)?;
let claims: Claims<T> = match token.content_type {
ContentType::Json => serde_json::from_slice(&token.serialized_claims)
.map_err(ValidationError::MalformedClaims)?,
ContentType::Cbor => serde_cbor::from_slice(&token.serialized_claims)
.map_err(ValidationError::MalformedCborClaims)?,
};
if !self.verify_signature(&signature, verifying_key, token.signed_data) {
return Err(ValidationError::InvalidSignature);
}
Ok(SignedToken {
signature,
token: Token {
header: token.header.clone(),
claims,
},
})
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct Header {
#[serde(rename = "jku", default, skip_serializing_if = "Option::is_none")]
pub key_set_url: Option<String>,
#[serde(rename = "kid", default, skip_serializing_if = "Option::is_none")]
pub key_id: Option<String>,
#[serde(rename = "x5u", default, skip_serializing_if = "Option::is_none")]
pub certificate_url: Option<String>,
#[serde(rename = "x5t", default, skip_serializing_if = "Option::is_none")]
pub certificate_thumbprint: Option<String>,
#[serde(rename = "typ", default, skip_serializing_if = "Option::is_none")]
pub signature_type: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct CompleteHeader<'a> {
#[serde(rename = "alg")]
algorithm: Cow<'a, str>,
#[serde(rename = "cty", default, skip_serializing_if = "Option::is_none")]
content_type: Option<String>,
#[serde(flatten)]
inner: Header,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ContentType {
Json,
Cbor,
}
#[derive(Debug, Clone)]
pub struct UntrustedToken<'a> {
signed_data: &'a [u8],
header: Header,
algorithm: String,
content_type: ContentType,
serialized_claims: Vec<u8>,
signature: SmallVec<[u8; SIGNATURE_SIZE]>,
}
#[derive(Debug, Clone)]
pub struct Token<T> {
header: Header,
claims: Claims<T>,
}
impl<T> Token<T> {
pub fn header(&self) -> &Header {
&self.header
}
pub fn claims(&self) -> &Claims<T> {
&self.claims
}
}
#[non_exhaustive]
pub struct SignedToken<A: Algorithm + ?Sized, T> {
pub signature: A::Signature,
pub token: Token<T>,
}
impl<A, T> fmt::Debug for SignedToken<A, T>
where
A: Algorithm,
A::Signature: fmt::Debug,
T: fmt::Debug,
{
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("SignedToken")
.field("token", &self.token)
.field("signature", &self.signature)
.finish()
}
}
impl<A, T> Clone for SignedToken<A, T>
where
A: Algorithm,
A::Signature: Clone,
T: Clone,
{
fn clone(&self) -> Self {
Self {
signature: self.signature.clone(),
token: self.token.clone(),
}
}
}
impl<'a> TryFrom<&'a str> for UntrustedToken<'a> {
type Error = ParseError;
fn try_from(s: &'a str) -> Result<Self, Self::Error> {
let token_parts: Vec<_> = s.splitn(4, '.').collect();
match &token_parts[..] {
[header, claims, signature] => {
let header = base64::decode_config(header, base64::URL_SAFE_NO_PAD)?;
let serialized_claims = base64::decode_config(claims, base64::URL_SAFE_NO_PAD)?;
let mut decoded_signature = smallvec![0; 3 * (signature.len() + 3) / 4];
let signature_len = base64::decode_config_slice(
signature,
base64::URL_SAFE_NO_PAD,
&mut decoded_signature[..],
)?;
decoded_signature.truncate(signature_len);
let header: CompleteHeader<'_> =
serde_json::from_slice(&header).map_err(ParseError::MalformedHeader)?;
let content_type = match header.content_type {
None => ContentType::Json,
Some(ref s) if s.eq_ignore_ascii_case("json") => ContentType::Json,
Some(ref s) if s.eq_ignore_ascii_case("cbor") => ContentType::Cbor,
Some(s) => return Err(ParseError::UnsupportedContentType(s)),
};
Ok(Self {
signed_data: s.rsplitn(2, '.').nth(1).unwrap().as_bytes(),
header: header.inner,
algorithm: header.algorithm.into_owned(),
content_type,
serialized_claims,
signature: decoded_signature,
})
}
_ => Err(ParseError::InvalidTokenStructure),
}
}
}
impl<'a> UntrustedToken<'a> {
pub fn header(&self) -> &Header {
&self.header
}
pub fn algorithm(&self) -> &str {
&self.algorithm
}
pub fn signature_bytes(&self) -> &[u8] {
&self.signature
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::alg::*;
use assert_matches::assert_matches;
type Obj = serde_json::Map<String, serde_json::Value>;
const HS256_TOKEN: &str = "eyJ0eXAiOiJKV1QiLA0KICJhbGciOiJIUzI1NiJ9.\
eyJpc3MiOiJqb2UiLA0KICJleHAiOjEzMDA4MTkzODAsDQogImh0dHA6Ly9leGFt\
cGxlLmNvbS9pc19yb290Ijp0cnVlfQ.\
dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
const HS256_KEY: &str = "AyM1SysPpbyDfgZld3umj1qzKObwVMkoqQ-EstJQLr_T-1qS0gZH75\
aKtMN3Yj0iPS4hcgUuTwjAzZr1Z9CAow";
#[test]
fn invalid_token_structure() {
let mangled_str = HS256_TOKEN.replace('.', "");
assert_matches!(
UntrustedToken::try_from(mangled_str.as_str()).unwrap_err(),
ParseError::InvalidTokenStructure
);
let mut mangled_str = HS256_TOKEN.to_owned();
let signature_start = mangled_str.rfind('.').unwrap();
mangled_str.truncate(signature_start);
assert_matches!(
UntrustedToken::try_from(mangled_str.as_str()).unwrap_err(),
ParseError::InvalidTokenStructure
);
let mut mangled_str = HS256_TOKEN.to_owned();
mangled_str.push('.');
assert_matches!(
UntrustedToken::try_from(mangled_str.as_str()).unwrap_err(),
ParseError::InvalidTokenStructure
);
}
#[test]
fn base64_error_during_parsing() {
let mangled_str = HS256_TOKEN.replace('0', "+");
assert_matches!(
UntrustedToken::try_from(mangled_str.as_str()).unwrap_err(),
ParseError::Base64(_)
);
let mut mangled_str = HS256_TOKEN.to_owned();
mangled_str.truncate(mangled_str.len() - 1);
assert_matches!(
UntrustedToken::try_from(mangled_str.as_str()).unwrap_err(),
ParseError::Base64(_)
);
}
#[test]
fn malformed_header() {
let mangled_headers = [
r#"{"alg":"HS256""#,
"{}",
r#"{"alg":5}"#,
r#"{"alg":[1,"foo"]}"#,
r#"{"alg":false}"#,
r#"{"alg":"HS256","alg":"none"}"#,
];
for mangled_header in &mangled_headers {
let mangled_header = base64::encode_config(mangled_header, base64::URL_SAFE_NO_PAD);
let mut mangled_str = HS256_TOKEN.to_owned();
mangled_str.replace_range(..mangled_str.find('.').unwrap(), &mangled_header);
assert_matches!(
UntrustedToken::try_from(mangled_str.as_str()).unwrap_err(),
ParseError::MalformedHeader(_)
);
}
}
#[test]
fn unsupported_content_type() {
let mangled_header = r#"{"alg":"HS256","cty":"txt"}"#;
let mangled_header = base64::encode_config(mangled_header, base64::URL_SAFE_NO_PAD);
let mut mangled_str = HS256_TOKEN.to_owned();
mangled_str.replace_range(..mangled_str.find('.').unwrap(), &mangled_header);
assert_matches!(
UntrustedToken::try_from(mangled_str.as_str()).unwrap_err(),
ParseError::UnsupportedContentType(ref s) if s == "txt"
);
}
#[test]
fn malformed_json_claims() {
let malformed_claims = [
r#"{"exp":1500000000"#,
r#"{"exp":"1500000000"}"#,
r#"{"exp":false}"#,
r#"{"exp":1500000000,"nbf":1400000000,"exp":1510000000}"#,
r#"{"exp":1500000000000000000000000000000000}"#,
];
let claims_start = HS256_TOKEN.find('.').unwrap() + 1;
let claims_end = HS256_TOKEN.rfind('.').unwrap();
let key = base64::decode_config(HS256_KEY, base64::URL_SAFE_NO_PAD).unwrap();
let key = Hs256Key::from(&*key);
for claims in &malformed_claims {
let encoded_claims = base64::encode_config(claims.as_bytes(), base64::URL_SAFE_NO_PAD);
let mut mangled_str = HS256_TOKEN.to_owned();
mangled_str.replace_range(claims_start..claims_end, &encoded_claims);
let token = UntrustedToken::try_from(mangled_str.as_str()).unwrap();
assert_matches!(
Hs256.validate_integrity::<Obj>(&token, &key).unwrap_err(),
ValidationError::MalformedClaims(_),
"Failing claims: {}",
claims
);
}
}
}