use std::fmt;
use thiserror::Error;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
#[non_exhaustive]
pub enum HexError {
#[error("invalid hex character")]
InvalidChar,
#[error("invalid hex length: {0}")]
InvalidLength(usize),
#[error("hex length mismatch: expected {expected} bytes, got {actual}")]
LengthMismatch {
expected: usize,
actual: usize,
},
#[error("hex decoder reported an internal overflow")]
Overflow,
}
impl From<faster_hex::Error> for HexError {
fn from(err: faster_hex::Error) -> Self {
match err {
faster_hex::Error::InvalidChar => Self::InvalidChar,
faster_hex::Error::InvalidLength(len) => Self::InvalidLength(len),
faster_hex::Error::Overflow => Self::Overflow,
}
}
}
#[must_use]
pub fn encode<T>(bytes: T) -> String
where
T: AsRef<[u8]>,
{
faster_hex::hex_string(bytes.as_ref())
}
pub fn encode_to_slice<T>(bytes: T, out: &mut [u8]) -> Result<(), HexError>
where
T: AsRef<[u8]>,
{
let bytes = bytes.as_ref();
let expected = bytes.len() * 2;
if out.len() != expected {
return Err(HexError::LengthMismatch {
expected,
actual: out.len(),
});
}
faster_hex::hex_encode(bytes, out).map_err(HexError::from)?;
Ok(())
}
pub fn decode<T>(input: T) -> Result<Vec<u8>, HexError>
where
T: AsRef<[u8]>,
{
let input = input.as_ref();
if input.len() % 2 != 0 {
return Err(HexError::InvalidLength(input.len()));
}
let mut out = vec![0_u8; input.len() / 2];
faster_hex::hex_decode(input, &mut out).map_err(HexError::from)?;
Ok(out)
}
pub fn decode_to_slice<T>(input: T, out: &mut [u8]) -> Result<(), HexError>
where
T: AsRef<[u8]>,
{
let input = input.as_ref();
let expected = out.len() * 2;
if input.len() != expected {
return Err(HexError::LengthMismatch {
expected,
actual: input.len(),
});
}
faster_hex::hex_decode(input, out).map_err(HexError::from)?;
Ok(())
}
pub fn fmt_lower<T>(bytes: T, f: &mut fmt::Formatter<'_>) -> fmt::Result
where
T: AsRef<[u8]>,
{
for byte in bytes.as_ref() {
write!(f, "{byte:02x}")?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use hex_literal::hex;
use super::*;
#[test]
fn encode_roundtrip() {
let bytes = hex!("deadbeef");
let encoded = encode(bytes);
assert_eq!(encoded, "deadbeef");
assert_eq!(decode(&encoded).unwrap(), bytes);
}
#[test]
fn encode_to_slice_exact() {
let bytes = hex!("00112233");
let mut buf = [0_u8; 8];
encode_to_slice(bytes, &mut buf).unwrap();
assert_eq!(&buf, b"00112233");
}
#[test]
fn encode_to_slice_wrong_length() {
let bytes = hex!("ab");
let mut buf = [0_u8; 4];
let err = encode_to_slice(bytes, &mut buf).unwrap_err();
assert!(matches!(
err,
HexError::LengthMismatch {
expected: 2,
actual: 4
}
));
}
#[test]
fn decode_to_slice_exact() {
let mut buf = [0_u8; 4];
decode_to_slice("deadbeef", &mut buf).unwrap();
assert_eq!(buf, hex!("deadbeef"));
}
#[test]
fn decode_odd_length() {
let err = decode("abc").unwrap_err();
assert!(matches!(err, HexError::InvalidLength(3)));
}
#[test]
fn decode_invalid_char() {
let err = decode("zz").unwrap_err();
assert_eq!(err, HexError::InvalidChar);
}
#[test]
fn decode_to_slice_length_mismatch() {
let mut buf = [0_u8; 2];
let err = decode_to_slice("aabbcc", &mut buf).unwrap_err();
assert!(matches!(
err,
HexError::LengthMismatch {
expected: 4,
actual: 6
}
));
}
#[test]
fn fmt_lower_matches_encode() {
struct Wrap([u8; 4]);
impl fmt::Display for Wrap {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt_lower(self.0, f)
}
}
assert_eq!(Wrap(hex!("0a1b2c3d")).to_string(), "0a1b2c3d");
}
}