use serde::{Deserialize, Serialize};
use std::fmt;
pub const ENVELOPE_LEN: usize = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum MessageClass {
Hello = 1,
MotionTelemetry = 2,
StatusTelemetry = 3,
BulkUplink = 4,
Ack = 5,
Heartbeat = 6,
ReliableCommand = 7,
Teleop = 8,
HeartbeatAck = 9,
Json = 10,
}
impl MessageClass {
#[inline]
pub fn as_u8(self) -> u8 {
self as u8
}
}
impl TryFrom<u8> for MessageClass {
type Error = EnvelopeError;
fn try_from(byte: u8) -> Result<Self, Self::Error> {
Ok(match byte {
1 => MessageClass::Hello,
2 => MessageClass::MotionTelemetry,
3 => MessageClass::StatusTelemetry,
4 => MessageClass::BulkUplink,
5 => MessageClass::Ack,
6 => MessageClass::Heartbeat,
7 => MessageClass::ReliableCommand,
8 => MessageClass::Teleop,
9 => MessageClass::HeartbeatAck,
10 => MessageClass::Json,
other => return Err(EnvelopeError::UnknownClass(other)),
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EnvelopeError {
EmptyFrame,
UnknownClass(u8),
}
impl fmt::Display for EnvelopeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
EnvelopeError::EmptyFrame => write!(f, "empty wire frame: missing message-class byte"),
EnvelopeError::UnknownClass(b) => write!(f, "unknown message-class byte: {b}"),
}
}
}
impl std::error::Error for EnvelopeError {}
pub fn encode(class: MessageClass, payload: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(ENVELOPE_LEN + payload.len());
buf.push(class.as_u8());
buf.extend_from_slice(payload);
buf
}
pub fn decode(frame: &[u8]) -> Result<(MessageClass, &[u8]), EnvelopeError> {
let (&class_byte, payload) = frame.split_first().ok_or(EnvelopeError::EmptyFrame)?;
let class = MessageClass::try_from(class_byte)?;
Ok((class, payload))
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct JsonMessage {
pub kind: String,
pub data: serde_json::Value,
}
impl JsonMessage {
pub fn new<T: Serialize>(kind: impl Into<String>, body: &T) -> Result<Self, serde_json::Error> {
Ok(Self {
kind: kind.into(),
data: serde_json::to_value(body)?,
})
}
pub fn encode(&self) -> Result<Vec<u8>, serde_json::Error> {
serde_json::to_vec(self)
}
pub fn decode(bytes: &[u8]) -> Result<Self, serde_json::Error> {
serde_json::from_slice(bytes)
}
pub fn parse<T: serde::de::DeserializeOwned>(&self) -> Result<T, serde_json::Error> {
serde_json::from_value(self.data.clone())
}
}
#[cfg(test)]
mod tests {
use super::*;
const ALL_CLASSES: &[MessageClass] = &[
MessageClass::Hello,
MessageClass::MotionTelemetry,
MessageClass::StatusTelemetry,
MessageClass::BulkUplink,
MessageClass::Ack,
MessageClass::Heartbeat,
MessageClass::ReliableCommand,
MessageClass::Teleop,
MessageClass::HeartbeatAck,
MessageClass::Json,
];
#[test]
fn envelope_len_is_one() {
assert_eq!(ENVELOPE_LEN, 1);
}
#[test]
fn class_byte_roundtrips() {
for &class in ALL_CLASSES {
let byte = class.as_u8();
assert_eq!(MessageClass::try_from(byte), Ok(class));
}
}
#[test]
fn class_list_is_exhaustive() {
let max = ALL_CLASSES.iter().map(|c| c.as_u8()).max().unwrap();
assert_eq!(
ALL_CLASSES.len() as u8,
max,
"ALL_CLASSES must list every variant 1..=max"
);
}
#[test]
fn zero_byte_is_unknown() {
assert_eq!(
MessageClass::try_from(0),
Err(EnvelopeError::UnknownClass(0))
);
}
#[test]
fn unknown_byte_is_rejected() {
assert_eq!(
MessageClass::try_from(255),
Err(EnvelopeError::UnknownClass(255))
);
}
#[test]
fn encode_decode_roundtrip_preserves_class_and_payload() {
let payload = b"capnp-bytes-here";
for &class in ALL_CLASSES {
let frame = encode(class, payload);
assert_eq!(frame.len(), ENVELOPE_LEN + payload.len());
let (decoded_class, decoded_payload) = decode(&frame).expect("decode");
assert_eq!(decoded_class, class);
assert_eq!(decoded_payload, payload);
}
}
#[test]
fn encode_decode_roundtrip_empty_payload() {
let frame = encode(MessageClass::Heartbeat, &[]);
let (class, payload) = decode(&frame).expect("decode");
assert_eq!(class, MessageClass::Heartbeat);
assert!(payload.is_empty());
}
#[test]
fn decode_empty_frame_errors() {
assert_eq!(decode(&[]), Err(EnvelopeError::EmptyFrame));
}
#[test]
fn decode_unknown_class_errors() {
assert_eq!(
decode(&[200, 1, 2, 3]),
Err(EnvelopeError::UnknownClass(200))
);
}
#[test]
fn json_message_roundtrips_through_envelope() {
#[derive(Debug, PartialEq, serde::Serialize, serde::Deserialize)]
struct SpawnBody {
swarm_id: String,
count: u32,
}
let body = SpawnBody {
swarm_id: "sw-1".into(),
count: 3,
};
let msg = JsonMessage::new("sim.spawn", &body).expect("wrap");
assert_eq!(msg.kind, "sim.spawn");
let frame = encode(MessageClass::Json, &msg.encode().unwrap());
let (class, payload) = decode(&frame).unwrap();
assert_eq!(class, MessageClass::Json);
let back = JsonMessage::decode(payload).unwrap();
assert_eq!(back, msg);
assert_eq!(back.parse::<SpawnBody>().unwrap(), body);
}
#[test]
fn manual_stream_framing_matches_encode() {
let payload = b"xyz";
let mut manual = vec![MessageClass::Ack.as_u8()];
manual.extend_from_slice(payload);
assert_eq!(manual, encode(MessageClass::Ack, payload));
}
}