use alloc::string::String;
use alloc::vec::Vec;
use core::fmt;
use crate::State;
use crate::adapters::bytes::decode_base64;
use crate::adapters::{Base64, BytesEncoding};
use crate::error::Error;
#[derive(Clone, Copy)]
pub struct BytesFormat(Repr);
#[derive(Clone, Copy)]
enum Repr {
Encoded {
name: &'static str,
encode: fn(&[u8], &mut String),
decode: fn(&str) -> Result<Vec<u8>, Error>,
},
Seq,
}
impl BytesFormat {
pub const BASE64: BytesFormat = BytesFormat::encoded::<Base64>();
pub const SEQ: BytesFormat = BytesFormat(Repr::Seq);
pub const fn encoded<E: BytesEncoding>() -> BytesFormat {
BytesFormat(Repr::Encoded {
name: E::NAME,
encode: E::encode,
decode: E::decode,
})
}
#[inline]
pub fn of(state: &State) -> BytesFormat {
state.get::<BytesFormat>().copied().unwrap_or_default()
}
#[inline]
pub fn set(self, state: &mut State) {
*state.get_mut::<BytesFormat>() = self;
}
pub fn name(&self) -> &'static str {
match self.0 {
Repr::Encoded { name, .. } => name,
Repr::Seq => "seq",
}
}
pub fn is_seq(&self) -> bool {
matches!(self.0, Repr::Seq)
}
pub fn encode(&self, bytes: &[u8]) -> Option<String> {
match self.0 {
Repr::Encoded { encode, .. } => {
let mut rv = String::new();
encode(bytes, &mut rv);
Some(rv)
}
Repr::Seq => None,
}
}
pub fn decode(&self, s: &str) -> Result<Vec<u8>, Error> {
match self.0 {
Repr::Encoded { decode, .. } => decode(s),
Repr::Seq => decode_base64(s),
}
}
}
impl Default for BytesFormat {
fn default() -> BytesFormat {
BytesFormat::BASE64
}
}
impl fmt::Debug for BytesFormat {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("BytesFormat").field(&self.name()).finish()
}
}
impl PartialEq for BytesFormat {
fn eq(&self, other: &Self) -> bool {
match (self.0, other.0) {
(Repr::Encoded { name: a, .. }, Repr::Encoded { name: b, .. }) => a == b,
(Repr::Seq, Repr::Seq) => true,
_ => false,
}
}
}
impl Eq for BytesFormat {}
#[cfg(test)]
mod tests {
use super::*;
use crate::adapters::Base64Url;
#[test]
fn test_format() {
assert_eq!(BytesFormat::default(), BytesFormat::BASE64);
assert_eq!(BytesFormat::encoded::<Base64>(), BytesFormat::BASE64);
assert_ne!(BytesFormat::encoded::<Base64Url>(), BytesFormat::BASE64);
assert_ne!(BytesFormat::SEQ, BytesFormat::BASE64);
assert_eq!(format!("{:?}", BytesFormat::SEQ), "BytesFormat(\"seq\")");
assert_eq!(BytesFormat::SEQ.decode("AQ==").unwrap(), b"\x01");
}
}