#[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 {
self.to_blake3(pepper)
}
#[must_use]
pub fn to_blake3(&self, pepper: &[u8]) -> TokenHash {
let mut hasher = blake3::Hasher::new();
hasher.update(pepper);
hasher.update(&self.0);
let hash = hasher.finalize();
TokenHash::Blake3V1(hash)
}
}
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>")
}
}
const ALG_BLAKE3_V1: u8 = 0;
#[derive(Clone)]
#[non_exhaustive]
pub enum TokenHash {
Blake3V1(blake3::Hash),
}
impl TokenHash {
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
}
}
}
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)),
}
}
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,
)?)?)
}
}
}