use crate::SerialError;
pub const DELIMITER: u8 = 0x00;
const MAX_RUN: usize = 254;
#[must_use]
pub const fn max_encoded_len(payload_len: usize) -> usize {
payload_len + payload_len / MAX_RUN + 2
}
pub fn encode(payload: &[u8], output: &mut [u8]) -> Result<usize, SerialError> {
if output.len() < 2 {
return Err(SerialError::BufferTooSmall);
}
let mut write = 1usize;
let mut code_index: Option<usize> = Some(0);
let mut code: u8 = 1;
let n = payload.len();
let mut i = 0usize;
while i < n {
let byte = payload[i];
i += 1;
if byte != DELIMITER {
write_at(output, write, byte)?;
write += 1;
code += 1;
if code == 0xFF {
set_code(output, code_index, 0xFF);
code = 1;
if i < n {
reserve(output, &mut write, &mut code_index)?;
} else {
code_index = None;
}
}
} else {
set_code(output, code_index, code);
code = 1;
reserve(output, &mut write, &mut code_index)?;
}
}
set_code(output, code_index, code);
write_at(output, write, DELIMITER)?;
write += 1;
Ok(write)
}
pub fn decode(frame: &[u8], output: &mut [u8]) -> Result<usize, SerialError> {
let mut write = 0usize;
let mut owed: u8 = 0;
let mut code: u8 = 0xFF;
for &byte in frame {
if byte == DELIMITER {
if owed != 0 {
return Err(SerialError::TruncatedFrame);
}
return Ok(write);
}
if owed != 0 {
write_at(output, write, byte)?;
write += 1;
owed -= 1;
} else {
if code != 0xFF {
write_at(output, write, DELIMITER)?;
write += 1;
}
code = byte;
owed = byte - 1;
}
}
if owed != 0 {
return Err(SerialError::TruncatedFrame);
}
Ok(write)
}
#[derive(Debug)]
pub struct CobsDecoder<const N: usize> {
buffer: [u8; N],
len: usize,
owed: u8,
code: u8,
complete: bool,
}
impl<const N: usize> CobsDecoder<N> {
#[must_use]
pub const fn new() -> Self {
Self {
buffer: [0u8; N],
len: 0,
owed: 0,
code: 0xFF,
complete: false,
}
}
pub fn reset(&mut self) {
self.len = 0;
self.owed = 0;
self.code = 0xFF;
self.complete = false;
}
pub fn push(&mut self, byte: u8) -> Result<Option<&[u8]>, SerialError> {
if self.complete {
self.reset();
}
if byte == DELIMITER {
if self.owed != 0 {
self.reset();
return Err(SerialError::TruncatedFrame);
}
self.complete = true;
return Ok(Some(&self.buffer[..self.len]));
}
if self.owed != 0 {
self.store(byte)?;
self.owed -= 1;
} else {
if self.code != 0xFF {
self.store(DELIMITER)?;
}
self.code = byte;
self.owed = byte - 1;
}
Ok(None)
}
fn store(&mut self, byte: u8) -> Result<(), SerialError> {
if self.len >= N {
self.reset();
return Err(SerialError::BufferTooSmall);
}
self.buffer[self.len] = byte;
self.len += 1;
Ok(())
}
}
impl<const N: usize> Default for CobsDecoder<N> {
fn default() -> Self {
Self::new()
}
}
fn write_at(output: &mut [u8], index: usize, byte: u8) -> Result<(), SerialError> {
if index >= output.len() {
return Err(SerialError::BufferTooSmall);
}
output[index] = byte;
Ok(())
}
fn set_code(output: &mut [u8], code_index: Option<usize>, code: u8) {
if let Some(index) = code_index {
output[index] = code;
}
}
fn reserve(
output: &[u8],
write: &mut usize,
code_index: &mut Option<usize>,
) -> Result<(), SerialError> {
if *write >= output.len() {
return Err(SerialError::BufferTooSmall);
}
*code_index = Some(*write);
*write += 1;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn range(start: u8, end: u8) -> Vec<u8> {
(start..=end).collect()
}
fn canonical_vectors() -> Vec<(Vec<u8>, Vec<u8>)> {
let mut v: Vec<(Vec<u8>, Vec<u8>)> = vec![
(vec![0x00], vec![0x01, 0x01, 0x00]),
(vec![0x00, 0x00], vec![0x01, 0x01, 0x01, 0x00]),
(vec![0x00, 0x11, 0x00], vec![0x01, 0x02, 0x11, 0x01, 0x00]),
(
vec![0x11, 0x22, 0x00, 0x33],
vec![0x03, 0x11, 0x22, 0x02, 0x33, 0x00],
),
(
vec![0x11, 0x22, 0x33, 0x44],
vec![0x05, 0x11, 0x22, 0x33, 0x44, 0x00],
),
(
vec![0x11, 0x00, 0x00, 0x00],
vec![0x02, 0x11, 0x01, 0x01, 0x01, 0x00],
),
];
let mut e7 = vec![0xFF];
e7.extend(range(0x01, 0xFE));
e7.push(0x00);
v.push((range(0x01, 0xFE), e7));
let mut p8 = vec![0x00];
p8.extend(range(0x01, 0xFE));
let mut e8 = vec![0x01, 0xFF];
e8.extend(range(0x01, 0xFE));
e8.push(0x00);
v.push((p8, e8));
let mut e9 = vec![0xFF];
e9.extend(range(0x01, 0xFE));
e9.extend([0x02, 0xFF, 0x00]);
v.push((range(0x01, 0xFF), e9));
let mut p10 = range(0x02, 0xFF);
p10.push(0x00);
let mut e10 = vec![0xFF];
e10.extend(range(0x02, 0xFF));
e10.extend([0x01, 0x01, 0x00]);
v.push((p10, e10));
let mut p11 = range(0x03, 0xFF);
p11.extend([0x00, 0x01]);
let mut e11 = vec![0xFE];
e11.extend(range(0x03, 0xFF));
e11.extend([0x02, 0x01, 0x00]);
v.push((p11, e11));
v
}
#[test]
fn encode_matches_the_canonical_vectors() {
for (payload, expected) in canonical_vectors() {
let mut out = vec![0u8; max_encoded_len(payload.len())];
let n = encode(&payload, &mut out).unwrap();
assert_eq!(&out[..n], &expected[..], "payload {payload:02x?}");
}
}
#[test]
fn decode_matches_the_canonical_vectors() {
for (payload, encoded) in canonical_vectors() {
let mut out = vec![0u8; payload.len()];
let n = decode(&encoded, &mut out).unwrap();
assert_eq!(&out[..n], &payload[..], "encoded {encoded:02x?}");
}
}
#[test]
fn the_encoding_never_contains_an_interior_zero() {
for (payload, _) in canonical_vectors() {
let mut out = vec![0u8; max_encoded_len(payload.len())];
let n = encode(&payload, &mut out).unwrap();
assert!(out[..n - 1].iter().all(|&b| b != 0));
assert_eq!(out[n - 1], DELIMITER);
}
}
#[test]
fn empty_payload_round_trips() {
let mut frame = [0u8; 4];
let n = encode(&[], &mut frame).unwrap();
assert_eq!(&frame[..n], &[0x01, 0x00]);
let mut out = [0u8; 4];
let m = decode(&frame[..n], &mut out).unwrap();
assert_eq!(m, 0);
}
#[test]
fn round_trips_payloads_across_the_run_boundary() {
for len in [0usize, 1, 2, 253, 254, 255, 256, 509, 510, 511] {
for &fill in &[0x00u8, 0x41, 0xFF] {
let payload = vec![fill; len];
let mut frame = vec![0u8; max_encoded_len(len)];
let n = encode(&payload, &mut frame).unwrap();
let mut out = vec![0u8; len];
let m = decode(&frame[..n], &mut out).unwrap();
assert_eq!(&out[..m], &payload[..], "len {len} fill {fill:#04x}");
}
}
}
#[test]
fn round_trips_a_mixed_payload_with_scattered_zeros() {
let payload: Vec<u8> = (0..600u16).map(|i| (i % 7) as u8).collect();
let mut frame = vec![0u8; max_encoded_len(payload.len())];
let n = encode(&payload, &mut frame).unwrap();
let mut out = vec![0u8; payload.len()];
let m = decode(&frame[..n], &mut out).unwrap();
assert_eq!(&out[..m], &payload[..]);
}
#[test]
fn decode_tolerates_a_missing_trailing_delimiter() {
let mut out = [0u8; 4];
let n = decode(&[0x03, 0x11, 0x22, 0x02, 0x33], &mut out).unwrap();
assert_eq!(&out[..n], &[0x11, 0x22, 0x00, 0x33]);
}
#[test]
fn a_code_that_overruns_the_frame_is_truncated() {
let mut out = [0u8; 4];
assert_eq!(
decode(&[0x03, 0x11, 0x00], &mut out),
Err(SerialError::TruncatedFrame)
);
}
#[test]
fn encode_reports_a_full_buffer() {
let mut frame = [0u8; 3];
assert_eq!(
encode(&[0x11, 0x22, 0x33], &mut frame),
Err(SerialError::BufferTooSmall)
);
}
#[test]
fn decode_reports_a_full_buffer() {
let mut out = [0u8; 1];
assert_eq!(
decode(&[0x03, 0x11, 0x22, 0x00], &mut out),
Err(SerialError::BufferTooSmall)
);
}
#[test]
fn streaming_decoder_matches_the_canonical_vectors() {
for (payload, encoded) in canonical_vectors() {
let mut decoder: CobsDecoder<512> = CobsDecoder::new();
let mut produced = false;
for &byte in &encoded {
if let Some(frame) = decoder.push(byte).unwrap() {
assert_eq!(frame, &payload[..], "encoded {encoded:02x?}");
produced = true;
}
}
assert!(produced, "no frame for {encoded:02x?}");
}
}
#[test]
fn streaming_decoder_reports_overflow_then_recovers() {
let mut decoder: CobsDecoder<2> = CobsDecoder::new();
assert!(decoder.push(0x04).unwrap().is_none());
assert!(decoder.push(0x11).unwrap().is_none());
assert!(decoder.push(0x22).unwrap().is_none());
assert_eq!(decoder.push(0x33), Err(SerialError::BufferTooSmall));
assert!(decoder.push(0x03).unwrap().is_none());
assert!(decoder.push(0xAA).unwrap().is_none());
assert!(decoder.push(0xBB).unwrap().is_none());
assert_eq!(decoder.push(DELIMITER).unwrap(), Some(&[0xAA, 0xBB][..]));
}
}