use serde::Serialize;
use serde::de::DeserializeOwned;
use crate::abi::CodecId;
pub const MAX_DECODE_BODY_BYTES: usize = 16 * 1024 * 1024;
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 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"));
}
}