#[cfg(feature = "encryption")]
use crate::crypto::{CipherSuite, EncryptedSecret};
#[cfg(feature = "serde")]
use serde_big_array::BigArray;
use crate::{encoding::base64, random};
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[cfg(feature = "encryption")]
#[error("Crypto error: {0}")]
Crypto(#[from] crate::crypto::Error),
#[error("Random error: {0}")]
Random(#[from] crate::random::Error),
#[error("Length mismatch: expected = {expected}; provided = {provided}")]
LengthMismatch { expected: usize, provided: usize },
#[error("Invalid base64: {0}")]
Base64(#[from] base64::DecodeSliceError),
#[error("Hash algorithm is missing")]
HashAlgorithmMissing,
#[error("Hash algorithm is unknown or invalid: {0}")]
HashAlgorithmInvalid(u8),
}
pub type Result<T> = std::result::Result<T, Error>;
#[derive(zeroize::ZeroizeOnDrop, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[cfg_attr(feature = "serde", serde(transparent))]
pub struct Token<const LEN: usize>(
#[cfg_attr(feature = "serde", serde(with = "BigArray"))] [u8; LEN],
);
impl<const LEN: usize> Token<LEN> {
#[must_use]
pub fn random() -> Result<Self> {
let mut s = [0u8; LEN];
random::fill_bytes(&mut s)?;
Ok(Self(s))
}
#[must_use]
pub fn from_bytes(data: [u8; LEN]) -> Self {
Self(data)
}
#[must_use]
pub fn to_base64(&self) -> String {
base64::encode(&self.0)
}
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
#[must_use]
pub fn from_base64(data: &str) -> Result<Self> {
let mut s = [0u8; LEN];
let written = base64::decode_slice(data, &mut s)?;
if written == LEN {
Ok(Self(s))
} else {
Err(Error::LengthMismatch {
expected: LEN,
provided: written,
})
}
}
#[cfg(feature = "encryption")]
#[must_use]
pub fn encrypt(&self, key: &[u8], algo: CipherSuite) -> Result<EncryptedToken<LEN>> {
let enc = EncryptedSecret::encrypt(self.as_bytes(), key, algo)?;
Ok(EncryptedToken::from_secret(enc))
}
#[must_use]
pub fn to_default_hash(&self, pepper: &[u8]) -> TokenHash {
TokenHash::compute(self, pepper)
}
#[must_use]
pub fn to_blake3(&self, pepper: &[u8]) -> TokenHash {
TokenHash::compute_blake3_v1(self, pepper)
}
}
impl<const LEN: usize> std::fmt::Debug for Token<LEN> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "<redacted>")
}
}
impl<const LEN: usize> std::fmt::Display for Token<LEN> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "<redacted>")
}
}
#[must_use]
fn compute_blake3_v1_raw(bytes: &[u8], pepper: &[u8]) -> blake3::Hash {
let mut hasher = blake3::Hasher::new();
hasher.update(pepper);
hasher.update(bytes);
hasher.finalize()
}
const ALG_BLAKE3_V1: u8 = 0;
#[derive(Clone, Hash)]
#[non_exhaustive]
pub enum TokenHash {
Blake3V1(blake3::Hash),
}
impl TokenHash {
#[must_use]
pub fn compute<const LEN: usize>(token: &Token<LEN>, pepper: &[u8]) -> Self {
Self::compute_blake3_v1(token, pepper)
}
#[must_use]
pub fn compute_blake3_v1<const LEN: usize>(token: &Token<LEN>, pepper: &[u8]) -> Self {
let hash = compute_blake3_v1_raw(token.as_bytes(), pepper);
Self::Blake3V1(hash)
}
#[must_use]
pub fn verify<const LEN: usize>(&self, token: &Token<LEN>, pepper: &[u8]) -> bool {
match self {
Self::Blake3V1(a) => {
let b = compute_blake3_v1_raw(token.as_bytes(), pepper);
*a == b
}
}
}
pub fn to_bytes(&self) -> Vec<u8> {
match &self {
Self::Blake3V1(h) => {
let mut res = Vec::with_capacity(1 + blake3::OUT_LEN);
res.push(ALG_BLAKE3_V1);
res.extend_from_slice(h.as_bytes());
res
}
}
}
#[must_use]
pub fn from_bytes(data: &[u8]) -> Result<Self> {
if data.len() < 1 {
return Err(Error::HashAlgorithmMissing);
}
match data[0] {
ALG_BLAKE3_V1 => Ok(Self::blake3_from_bytes(&data[1..])?),
algo => Err(Error::HashAlgorithmInvalid(algo)),
}
}
#[must_use]
fn blake3_from_bytes(data: &[u8]) -> Result<Self> {
let hash = blake3::Hash::from_slice(data).map_err(|_| Error::LengthMismatch {
expected: blake3::OUT_LEN,
provided: data.len(),
})?;
Ok(Self::Blake3V1(hash))
}
}
impl std::fmt::Debug for TokenHash {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "<redacted>")
}
}
impl std::fmt::Display for TokenHash {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "<redacted>")
}
}
#[cfg(feature = "serde")]
mod serde_impl {
use super::*;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
impl Serialize for TokenHash {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: Serializer,
{
let bytes = self.to_bytes();
serializer.serialize_bytes(&bytes)
}
}
impl<'de> Deserialize<'de> for TokenHash {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let bytes = <Vec<u8>>::deserialize(deserializer)?;
TokenHash::from_bytes(&bytes).map_err(serde::de::Error::custom)
}
}
}
#[cfg(feature = "encryption")]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(transparent))]
#[cfg_attr(feature = "sqlx", derive(sqlx::Type))]
#[cfg_attr(feature = "sqlx", sqlx(transparent))]
#[derive(Clone, Debug)]
pub struct EncryptedToken<const LEN: usize>(EncryptedSecret);
#[cfg(feature = "encryption")]
impl<const LEN: usize> EncryptedToken<LEN> {
#[must_use]
pub fn from_encrypted_bytes(input: &[u8]) -> Result<Self> {
let secret = EncryptedSecret::from_encrypted_bytes(input)?;
Ok(Self(secret))
}
#[must_use]
pub fn from_secret(secret: EncryptedSecret) -> Self {
Self(secret)
}
#[must_use]
pub fn to_encrypted_bytes(&self) -> Vec<u8> {
self.0.to_encrypted_bytes()
}
#[must_use]
pub fn decrypt(&self, key: &[u8]) -> Result<Token<LEN>> {
let plaintext = self.0.decrypt(key)?;
let provided_len = plaintext.len();
let arr = plaintext.try_into().map_err(|_| Error::LengthMismatch {
expected: LEN,
provided: provided_len,
})?;
Ok(Token::from_bytes(arr))
}
pub fn algorithm(&self) -> CipherSuite {
self.0.algorithm()
}
}
#[cfg(feature = "encryption")]
impl<const LEN: usize> std::fmt::Display for EncryptedToken<LEN> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.0)
}
}
#[cfg(feature = "sqlx")]
mod sqlx_impl {
use super::*;
use sqlx::{Database, Decode, Encode, Type, encode::IsNull, error::BoxDynError};
#[cfg(feature = "sqlx_postgres")]
impl sqlx::postgres::PgHasArrayType for TokenHash {
fn array_type_info() -> sqlx::postgres::PgTypeInfo {
<Vec<u8> as sqlx::postgres::PgHasArrayType>::array_type_info()
}
}
impl<'a, DB: Database> Type<DB> for TokenHash
where
&'a [u8]: Type<DB>,
{
fn type_info() -> <DB as Database>::TypeInfo {
<&[u8] as Type<DB>>::type_info()
}
fn compatible(ty: &<DB as Database>::TypeInfo) -> bool {
<&[u8] as Type<DB>>::compatible(ty)
}
}
impl<'a, DB: Database> Encode<'a, DB> for TokenHash
where
Vec<u8>: Encode<'a, DB>,
{
fn encode_by_ref(
&self,
buf: &mut <DB as Database>::ArgumentBuffer<'a>,
) -> std::result::Result<IsNull, BoxDynError> {
<Vec<u8> as Encode<'a, DB>>::encode_by_ref(&self.to_bytes(), buf)
}
fn produces(&self) -> Option<<DB as Database>::TypeInfo> {
<Vec<u8> as Encode<'a, DB>>::produces(&self.to_bytes())
}
fn size_hint(&self) -> usize {
<Vec<u8> as Encode<'a, DB>>::size_hint(&self.to_bytes())
}
}
impl<'a, DB: Database> Decode<'a, DB> for TokenHash
where
Vec<u8>: Decode<'a, DB>,
{
fn decode(value: <DB as Database>::ValueRef<'a>) -> std::result::Result<Self, BoxDynError> {
Ok(Self::from_bytes(&<Vec<u8> as Decode<'a, DB>>::decode(
value,
)?)?)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hash_verify_success() {
let token = Token::<32>::random().unwrap();
let pepper = b"secret-pepper";
let hash = token.to_default_hash(pepper);
assert!(hash.verify(&token, pepper));
}
#[test]
fn hash_verify_wrong_pepper_fails() {
let token = Token::<32>::random().unwrap();
let hash = token.to_default_hash(b"pepper1");
assert!(!hash.verify(&token, b"pepper2"));
}
#[test]
fn hash_verify_wrong_token_fails() {
let token1 = Token::<32>::random().unwrap();
let token2 = Token::<32>::random().unwrap();
let hash = token1.to_default_hash(b"pepper");
assert!(!hash.verify(&token2, b"pepper"));
}
#[test]
fn token_hash_roundtrip_bytes() {
let token = Token::<32>::random().unwrap();
let hash = token.to_default_hash(b"pepper");
let bytes = hash.to_bytes();
let decoded = TokenHash::from_bytes(&bytes).unwrap();
assert!(decoded.verify(&token, b"pepper"));
}
#[test]
fn token_hash_rejects_invalid_length() {
let data = vec![0, 1, 2];
assert!(TokenHash::from_bytes(&data).is_err());
}
#[test]
fn token_hash_rejects_unknown_algorithm() {
let mut data = vec![255]; data.extend([0u8; 32]);
assert!(matches!(
TokenHash::from_bytes(&data),
Err(Error::HashAlgorithmInvalid(255))
));
}
#[test]
fn base64_roundtrip() {
let token = Token::<32>::random().unwrap();
let encoded = token.to_base64();
let decoded = Token::<32>::from_base64(&encoded).unwrap();
assert_eq!(token.as_bytes(), decoded.as_bytes());
}
#[test]
fn base64_rejects_wrong_length() {
let token = Token::<32>::random().unwrap();
let mut encoded = token.to_base64();
encoded.pop();
assert!(Token::<32>::from_base64(&encoded).is_err());
}
}