pub mod algorithm;
pub mod errors;
pub mod hkdf;
use base64::{engine::general_purpose, Engine as _};
use serde::{Deserialize, Serialize};
pub use algorithm::Algorithm;
pub use errors::Error;
#[cfg(not(feature = "ring"))]
use hmac::Mac;
#[cfg(feature = "ring")]
use ring::hmac;
pub const DELIM: char = '.';
#[derive(Default, Debug, Clone)]
pub enum Encoder {
Standard,
UrlSafe,
StandardNoPadding,
#[default]
UrlSafeNoPadding,
}
impl Encoder {
#[inline]
fn get_encoder(&self) -> general_purpose::GeneralPurpose {
match self {
Encoder::Standard => general_purpose::STANDARD,
Encoder::UrlSafe => general_purpose::URL_SAFE,
Encoder::StandardNoPadding => general_purpose::STANDARD_NO_PAD,
Encoder::UrlSafeNoPadding => general_purpose::URL_SAFE_NO_PAD,
}
}
}
pub trait Payload {
fn get_exp(&self) -> Option<chrono::DateTime<chrono::Utc>>;
}
#[derive(Debug, Clone)]
pub struct KeyInfo {
pub key: Vec<u8>,
pub salt: Vec<u8>,
pub info: Vec<u8>,
}
impl Default for KeyInfo {
fn default() -> Self {
Self {
key: vec![],
salt: vec![],
info: vec![],
}
}
}
#[derive(Debug, Clone)]
pub struct HmacSigner {
#[cfg(not(feature = "ring"))]
expanded_key: Vec<u8>,
#[cfg(not(feature = "ring"))]
algo: Algorithm,
#[cfg(feature = "ring")]
expanded_key: hmac::Key,
encoder: general_purpose::GeneralPurpose,
}
#[cfg(not(feature = "ring"))]
macro_rules! get_hmac {
($self:ident, $D:ty) => {
hmac::Hmac::<$D>::new_from_slice(&$self.expanded_key)
.expect("HMAC can take key of any size")
};
}
#[cfg(not(feature = "ring"))]
macro_rules! hmac_sign {
($self:ident, $payload:ident, $D:ty) => {{
let mut mac = get_hmac!($self, $D);
mac.update($payload);
mac.finalize().into_bytes().to_vec()
}};
}
#[cfg(not(feature = "ring"))]
macro_rules! hmac_verify {
($self:ident, $payload:ident, $signature:ident, $D:ty) => {{
let mut mac = get_hmac!($self, $D);
mac.update($payload);
mac.verify_slice($signature).is_ok()
}};
}
impl HmacSigner {
pub fn new(key_info: KeyInfo, algo: Algorithm, encoder: Encoder) -> Self {
if key_info.key.is_empty() {
panic!("Key cannot be empty"); }
let expanded_key = hkdf::HkdfWrapper::new(algo.clone()).expand(
&key_info.key,
&key_info.salt,
&key_info.info,
);
#[cfg(feature = "ring")]
{
let expanded_key = hmac::Key::new(algo.to_hmac(), &expanded_key);
return Self {
expanded_key,
encoder: encoder.get_encoder(),
};
}
#[cfg(not(feature = "ring"))]
Self {
expanded_key,
algo,
encoder: encoder.get_encoder(),
}
}
#[inline]
#[cfg(not(feature = "ring"))]
fn sign_payload(&self, payload: &[u8]) -> Vec<u8> {
match self.algo {
Algorithm::SHA1 => hmac_sign!(self, payload, sha1::Sha1),
Algorithm::SHA256 => hmac_sign!(self, payload, sha2::Sha256),
Algorithm::SHA384 => hmac_sign!(self, payload, sha2::Sha384),
Algorithm::SHA512 => hmac_sign!(self, payload, sha2::Sha512),
}
}
#[inline]
#[cfg(not(feature = "ring"))]
fn verify(&self, payload: &[u8], signature: &[u8]) -> bool {
match self.algo {
Algorithm::SHA1 => hmac_verify!(self, payload, signature, sha1::Sha1),
Algorithm::SHA256 => hmac_verify!(self, payload, signature, sha2::Sha256),
Algorithm::SHA384 => hmac_verify!(self, payload, signature, sha2::Sha384),
Algorithm::SHA512 => hmac_verify!(self, payload, signature, sha2::Sha512),
}
}
#[inline]
#[cfg(feature = "ring")]
fn sign_payload(&self, payload: &[u8]) -> Vec<u8> {
hmac::sign(&self.expanded_key, payload).as_ref().to_vec()
}
#[inline]
#[cfg(feature = "ring")]
fn verify(&self, payload: &[u8], signature: &[u8]) -> bool {
hmac::verify(&self.expanded_key, payload, signature).is_ok()
}
}
impl HmacSigner {
pub fn unsign<T: for<'de> Deserialize<'de> + Payload>(&self, token: &str) -> Result<T, Error> {
let parts: Vec<&str> = token.split(DELIM).collect();
if parts.len() != 2 {
return Err(Error::InvalidInput(token.to_string()));
}
let encoded_payload = parts[0];
if encoded_payload.is_empty() {
return Err(Error::InvalidToken);
}
let signature = self
.encoder
.decode(parts[1])
.map_err(|_| Error::InvalidSignature)?;
let encoded_payload = parts[0].as_bytes();
if !self.verify(&encoded_payload, &signature) {
return Err(Error::InvalidToken);
}
let decoded_payload = self
.encoder
.decode(encoded_payload)
.expect("payload should be valid base64");
let payload = String::from_utf8(decoded_payload).expect("payload should be valid utf-8");
let deserialised_payload: T =
serde_json::from_str(&payload).map_err(|_| Error::InvalidPayload)?;
if let Some(expiry) = deserialised_payload.get_exp() {
if expiry < chrono::Utc::now() {
return Err(Error::TokenExpired);
}
}
Ok(deserialised_payload)
}
pub fn sign<T: Serialize + Payload>(&self, payload: &T) -> String {
let token = serde_json::to_string(payload).unwrap();
let token = self.encoder.encode(token.as_bytes());
let signature = self.sign_payload(token.as_bytes());
let signature = self.encoder.encode(&signature);
format!("{}{}{}", token, DELIM, signature)
}
}