use std::str::FromStr;
use serde::Serialize;
use serde::de::DeserializeOwned;
pub const MAX_DECODE_BODY_BYTES: usize = 16 * 1024 * 1024;
const ENCODING_PREFIX: &str = "phoxal/v0";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[repr(u8)]
pub enum CodecId {
MessagePack = 1,
}
impl CodecId {
pub const fn as_u8(self) -> u8 {
self as u8
}
pub fn from_u8(value: u8) -> Option<Self> {
match value {
1 => Some(CodecId::MessagePack),
_ => None,
}
}
pub fn encoding_string(self) -> String {
format!("{ENCODING_PREFIX};codec={}", self.as_u8())
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct EncodingMetadata {
pub codec: u8,
}
impl EncodingMetadata {
pub fn codec_id(&self) -> Option<CodecId> {
CodecId::from_u8(self.codec)
}
}
impl FromStr for EncodingMetadata {
type Err = EncodingError;
fn from_str(value: &str) -> Result<Self, EncodingError> {
let mut parts = value.split(';');
let prefix = parts.next().unwrap_or_default();
if prefix != ENCODING_PREFIX {
return Err(EncodingError::Prefix {
found: prefix.to_string(),
});
}
let mut codec = None;
for field in parts {
let (key, value) =
field
.split_once('=')
.ok_or_else(|| EncodingError::MissingAssignment {
field: field.to_string(),
})?;
if value.is_empty() {
return Err(EncodingError::EmptyField {
field: key.to_string(),
});
}
match key {
"codec" => {
let parsed = value.parse::<u8>().map_err(|_| EncodingError::NotAU8 {
field: key.to_string(),
value: value.to_string(),
})?;
if codec.replace(parsed).is_some() {
return Err(EncodingError::DuplicateField {
field: key.to_string(),
});
}
}
_ => {
return Err(EncodingError::UnknownField {
field: key.to_string(),
});
}
}
}
Ok(EncodingMetadata {
codec: codec.ok_or(EncodingError::MissingCodec)?,
})
}
}
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum EncodingError {
#[error("expected encoding prefix '{ENCODING_PREFIX}', got '{found}'")]
Prefix {
found: String,
},
#[error("encoding field '{field}' is missing '='")]
MissingAssignment {
field: String,
},
#[error("encoding field '{field}' is empty")]
EmptyField {
field: String,
},
#[error("encoding field '{field}' is not a u8: '{value}'")]
NotAU8 {
field: String,
value: String,
},
#[error("duplicate encoding field '{field}'")]
DuplicateField {
field: String,
},
#[error("unknown encoding field '{field}'")]
UnknownField {
field: String,
},
#[error("encoding string is missing codec")]
MissingCodec,
}
pub(crate) fn truncate_utf8(value: &str, max_bytes: usize) -> String {
if value.len() <= max_bytes {
return value.to_string();
}
let mut end = max_bytes;
while !value.is_char_boundary(end) {
end -= 1;
}
value[..end].to_string()
}
pub trait Codec {
const ID: CodecId;
fn encode<T: Serialize>(value: &T) -> Result<Vec<u8>, CodecError>;
fn decode<T: DeserializeOwned>(bytes: &[u8]) -> Result<T, CodecError>;
}
pub struct MessagePack;
impl Codec for MessagePack {
const ID: CodecId = CodecId::MessagePack;
fn encode<T: Serialize>(value: &T) -> Result<Vec<u8>, CodecError> {
rmp_serde::to_vec_named(value).map_err(|e| CodecError::Encode(e.to_string()))
}
fn decode<T: DeserializeOwned>(bytes: &[u8]) -> Result<T, CodecError> {
if bytes.len() > MAX_DECODE_BODY_BYTES {
return Err(CodecError::Decode(format!(
"body is {} bytes; maximum is {MAX_DECODE_BODY_BYTES}",
bytes.len()
)));
}
rmp_serde::from_slice(bytes).map_err(|e| CodecError::Decode(e.to_string()))
}
}
#[derive(Debug, thiserror::Error)]
pub enum CodecError {
#[error("failed to encode body: {0}")]
Encode(String),
#[error("failed to decode body: {0}")]
Decode(String),
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn an_encoding_string_carries_only_the_codec() {
let encoding = CodecId::MessagePack.encoding_string();
assert_eq!(encoding, "phoxal/v0;codec=1");
let parsed: EncodingMetadata = encoding.parse().expect("the codec string parses");
assert_eq!(parsed.codec_id(), Some(CodecId::MessagePack));
}
#[test]
fn an_unknown_codec_id_parses_but_does_not_resolve() {
let parsed: EncodingMetadata = "phoxal/v0;codec=99".parse().expect("well-formed");
assert_eq!(parsed.codec, 99);
assert_eq!(parsed.codec_id(), None);
}
#[test]
fn a_malformed_encoding_string_names_what_was_wrong() {
assert_eq!(
"other/v0;codec=1".parse::<EncodingMetadata>(),
Err(EncodingError::Prefix {
found: "other/v0".to_string()
})
);
assert_eq!(
"".parse::<EncodingMetadata>(),
Err(EncodingError::Prefix {
found: String::new()
})
);
assert_eq!(
"phoxal/v0".parse::<EncodingMetadata>(),
Err(EncodingError::MissingCodec)
);
assert_eq!(
"phoxal/v0;codec".parse::<EncodingMetadata>(),
Err(EncodingError::MissingAssignment {
field: "codec".to_string()
})
);
assert_eq!(
"phoxal/v0;codec=".parse::<EncodingMetadata>(),
Err(EncodingError::EmptyField {
field: "codec".to_string()
})
);
assert_eq!(
"phoxal/v0;codec=x".parse::<EncodingMetadata>(),
Err(EncodingError::NotAU8 {
field: "codec".to_string(),
value: "x".to_string()
})
);
assert_eq!(
"phoxal/v0;codec=1;codec=1".parse::<EncodingMetadata>(),
Err(EncodingError::DuplicateField {
field: "codec".to_string()
})
);
assert_eq!(
"phoxal/v0;codec=1;schema=7".parse::<EncodingMetadata>(),
Err(EncodingError::UnknownField {
field: "schema".to_string()
})
);
}
#[test]
fn the_bootstrap_encoding_and_codec_are_pinned_to_their_literals() {
assert_eq!(ENCODING_PREFIX, "phoxal/v0");
assert_eq!(CodecId::MessagePack.as_u8(), 1);
assert_eq!(CodecId::MessagePack.encoding_string(), "phoxal/v0;codec=1");
assert_eq!(MessagePack::ID, CodecId::MessagePack);
#[derive(serde::Serialize)]
struct Body {
schema: &'static str,
}
let encoded = MessagePack::encode(&Body { schema: "v0" }).expect("the body encodes");
assert_eq!(encoded[0], 0x81, "codec 1 writes a map, not an array");
assert_eq!(&encoded[1..8], b"\xa6schema");
}
#[test]
fn messagepack_rejects_oversized_bodies_before_deserializing() {
let bytes = vec![0_u8; MAX_DECODE_BODY_BYTES + 1];
let error = MessagePack::decode::<Vec<u8>>(&bytes).expect_err("oversized body must fail");
assert!(error.to_string().contains("maximum"));
}
}