use core::fmt;
use sha2::{Digest, Sha256};
pub const COMMAND_ENVELOPE_FORMAT_VERSION: u32 = 1;
pub const MIN_SUPPORTED_FORMAT_VERSION: u32 = 1;
pub const MAX_COMMAND_PAYLOAD_BYTES: usize = 64 * 1024 * 1024;
pub const HEADER_LEN: usize = 4 + 16 + 4 + 4;
pub const CHECKSUM_LEN: usize = 32;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum EnvelopeError {
#[error("unsupported command-envelope format version {found} (supported {min}..={max})")]
UnsupportedVersion {
found: u32,
min: u32,
max: u32,
},
#[error("command-envelope checksum mismatch")]
ChecksumMismatch,
#[error("command envelope truncated: expected at least {expected} bytes, got {actual}")]
Truncated {
expected: usize,
actual: usize,
},
#[error("command envelope has {0} trailing bytes")]
TrailingBytes(usize),
#[error("command payload too large: {0} bytes")]
PayloadTooLarge(usize),
}
#[derive(Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct CommandEnvelope {
pub format_version: u32,
pub command_id: [u8; 16],
pub command_type: u32,
pub payload: Vec<u8>,
pub payload_sha256: [u8; 32],
}
impl CommandEnvelope {
pub fn new(command_type: u32, command_id: [u8; 16], payload: Vec<u8>) -> Self {
let payload_sha256 =
Self::checksum(COMMAND_ENVELOPE_FORMAT_VERSION, command_type, &payload);
Self {
format_version: COMMAND_ENVELOPE_FORMAT_VERSION,
command_id,
command_type,
payload,
payload_sha256,
}
}
pub fn checksum(format_version: u32, command_type: u32, payload: &[u8]) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(format_version.to_le_bytes());
hasher.update(command_type.to_le_bytes());
hasher.update((payload.len() as u64).to_le_bytes());
hasher.update(payload);
hasher.finalize().into()
}
pub fn verify(&self) -> Result<(), EnvelopeError> {
if !(MIN_SUPPORTED_FORMAT_VERSION..=COMMAND_ENVELOPE_FORMAT_VERSION)
.contains(&self.format_version)
{
return Err(EnvelopeError::UnsupportedVersion {
found: self.format_version,
min: MIN_SUPPORTED_FORMAT_VERSION,
max: COMMAND_ENVELOPE_FORMAT_VERSION,
});
}
if self.payload.len() > MAX_COMMAND_PAYLOAD_BYTES {
return Err(EnvelopeError::PayloadTooLarge(self.payload.len()));
}
let expected = Self::checksum(self.format_version, self.command_type, &self.payload);
if expected != self.payload_sha256 {
return Err(EnvelopeError::ChecksumMismatch);
}
Ok(())
}
pub fn encode(&self) -> Vec<u8> {
debug_assert!(
self.payload.len() <= MAX_COMMAND_PAYLOAD_BYTES,
"payload exceeds MAX_COMMAND_PAYLOAD_BYTES"
);
let mut out = Vec::with_capacity(HEADER_LEN + self.payload.len() + CHECKSUM_LEN);
out.extend_from_slice(&self.format_version.to_le_bytes());
out.extend_from_slice(&self.command_id);
out.extend_from_slice(&self.command_type.to_le_bytes());
out.extend_from_slice(&(self.payload.len() as u32).to_le_bytes());
out.extend_from_slice(&self.payload);
out.extend_from_slice(&self.payload_sha256);
out
}
pub fn decode(bytes: &[u8]) -> Result<Self, EnvelopeError> {
if bytes.len() < HEADER_LEN + CHECKSUM_LEN {
return Err(EnvelopeError::Truncated {
expected: HEADER_LEN + CHECKSUM_LEN,
actual: bytes.len(),
});
}
let format_version = u32::from_le_bytes(bytes[0..4].try_into().expect("slice len"));
let command_id: [u8; 16] = bytes[4..20].try_into().expect("slice len");
let command_type = u32::from_le_bytes(bytes[20..24].try_into().expect("slice len"));
let payload_len = u32::from_le_bytes(bytes[24..28].try_into().expect("slice len")) as usize;
if payload_len > MAX_COMMAND_PAYLOAD_BYTES {
return Err(EnvelopeError::PayloadTooLarge(payload_len));
}
let expected_total = HEADER_LEN + payload_len + CHECKSUM_LEN;
if bytes.len() < expected_total {
return Err(EnvelopeError::Truncated {
expected: expected_total,
actual: bytes.len(),
});
}
if bytes.len() > expected_total {
return Err(EnvelopeError::TrailingBytes(bytes.len() - expected_total));
}
let payload = bytes[HEADER_LEN..HEADER_LEN + payload_len].to_vec();
let payload_sha256: [u8; 32] = bytes[HEADER_LEN + payload_len..expected_total]
.try_into()
.expect("slice len");
let envelope = Self {
format_version,
command_id,
command_type,
payload,
payload_sha256,
};
envelope.verify()?;
Ok(envelope)
}
}
impl fmt::Debug for CommandEnvelope {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CommandEnvelope")
.field("format_version", &self.format_version)
.field("command_id", &self.command_id)
.field("command_type", &self.command_type)
.field("payload_len", &self.payload.len())
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip() {
let envelope = CommandEnvelope::new(7, [42u8; 16], b"hello".to_vec());
let bytes = envelope.encode();
assert_eq!(CommandEnvelope::decode(&bytes).unwrap(), envelope);
}
#[test]
fn bit_flip_breaks_checksum() {
let envelope = CommandEnvelope::new(1, [1u8; 16], vec![9u8; 64]);
let mut bytes = envelope.encode();
bytes[HEADER_LEN] ^= 0x01;
assert_eq!(
CommandEnvelope::decode(&bytes),
Err(EnvelopeError::ChecksumMismatch)
);
}
#[test]
fn unknown_version_fails_closed() {
let mut envelope = CommandEnvelope::new(1, [1u8; 16], vec![]);
envelope.format_version = COMMAND_ENVELOPE_FORMAT_VERSION + 1;
envelope.payload_sha256 =
CommandEnvelope::checksum(envelope.format_version, 1, &envelope.payload);
assert!(matches!(
envelope.verify(),
Err(EnvelopeError::UnsupportedVersion { .. })
));
}
#[test]
fn truncation_and_trailing_bytes_fail() {
let envelope = CommandEnvelope::new(3, [7u8; 16], vec![1, 2, 3]);
let bytes = envelope.encode();
assert!(matches!(
CommandEnvelope::decode(&bytes[..bytes.len() - 1]),
Err(EnvelopeError::Truncated { .. })
));
let mut longer = bytes.clone();
longer.push(0);
assert_eq!(
CommandEnvelope::decode(&longer),
Err(EnvelopeError::TrailingBytes(1))
);
}
}