use crate::SerialError;
pub const END: u8 = 0xC0;
pub const ESC: u8 = 0xDB;
pub const ESC_END: u8 = 0xDC;
pub const ESC_ESC: u8 = 0xDD;
#[must_use]
pub const fn max_encoded_len(payload_len: usize) -> usize {
payload_len * 2 + 1
}
pub fn encode(payload: &[u8], output: &mut [u8]) -> Result<usize, SerialError> {
let mut write = 0usize;
for &byte in payload {
match byte {
END => {
push(output, &mut write, ESC)?;
push(output, &mut write, ESC_END)?;
}
ESC => {
push(output, &mut write, ESC)?;
push(output, &mut write, ESC_ESC)?;
}
other => push(output, &mut write, other)?,
}
}
push(output, &mut write, END)?;
Ok(write)
}
pub fn decode(frame: &[u8], output: &mut [u8]) -> Result<usize, SerialError> {
let mut write = 0usize;
let mut in_escape = false;
for &byte in frame {
if byte == END {
if in_escape {
return Err(SerialError::TruncatedFrame);
}
if write == 0 {
continue;
}
return Ok(write);
}
if in_escape {
let decoded = match byte {
ESC_END => END,
ESC_ESC => ESC,
_ => return Err(SerialError::InvalidEscape),
};
push(output, &mut write, decoded)?;
in_escape = false;
} else if byte == ESC {
in_escape = true;
} else {
push(output, &mut write, byte)?;
}
}
if in_escape {
return Err(SerialError::TruncatedFrame);
}
Ok(write)
}
#[derive(Debug)]
pub struct SlipDecoder<const N: usize> {
buffer: [u8; N],
len: usize,
in_escape: bool,
complete: bool,
}
impl<const N: usize> SlipDecoder<N> {
#[must_use]
pub const fn new() -> Self {
Self {
buffer: [0u8; N],
len: 0,
in_escape: false,
complete: false,
}
}
pub fn reset(&mut self) {
self.len = 0;
self.in_escape = false;
self.complete = false;
}
pub fn push(&mut self, byte: u8) -> Result<Option<&[u8]>, SerialError> {
if self.complete {
self.reset();
}
if byte == END {
if self.in_escape {
self.reset();
return Err(SerialError::TruncatedFrame);
}
if self.len == 0 {
return Ok(None);
}
self.complete = true;
return Ok(Some(&self.buffer[..self.len]));
}
if self.in_escape {
let decoded = match byte {
ESC_END => END,
ESC_ESC => ESC,
_ => {
self.reset();
return Err(SerialError::InvalidEscape);
}
};
self.in_escape = false;
self.store(decoded)?;
} else if byte == ESC {
self.in_escape = true;
} else {
self.store(byte)?;
}
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 SlipDecoder<N> {
fn default() -> Self {
Self::new()
}
}
fn push(output: &mut [u8], write: &mut usize, byte: u8) -> Result<(), SerialError> {
if *write >= output.len() {
return Err(SerialError::BufferTooSmall);
}
output[*write] = byte;
*write += 1;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn constants_match_rfc_1055() {
assert_eq!(END, 0o300);
assert_eq!(ESC, 0o333);
assert_eq!(ESC_END, 0o334);
assert_eq!(ESC_ESC, 0o335);
}
#[test]
fn plain_payload_just_gets_a_trailing_end() {
let mut frame = [0u8; 8];
let n = encode(b"hi", &mut frame).unwrap();
assert_eq!(&frame[..n], &[b'h', b'i', END]);
}
#[test]
fn a_literal_end_byte_is_escaped() {
let mut frame = [0u8; 8];
let n = encode(&[END], &mut frame).unwrap();
assert_eq!(&frame[..n], &[ESC, ESC_END, END]);
}
#[test]
fn a_literal_esc_byte_is_escaped() {
let mut frame = [0u8; 8];
let n = encode(&[ESC], &mut frame).unwrap();
assert_eq!(&frame[..n], &[ESC, ESC_ESC, END]);
}
#[test]
fn both_specials_in_one_payload() {
let mut frame = [0u8; 16];
let n = encode(&[END, ESC, 0x01], &mut frame).unwrap();
assert_eq!(&frame[..n], &[ESC, ESC_END, ESC, ESC_ESC, 0x01, END]);
}
#[test]
fn round_trips_every_byte_value() {
let payload: [u8; 256] = core::array::from_fn(|i| i as u8);
let mut frame = [0u8; max_encoded_len(256)];
let n = encode(&payload, &mut frame).unwrap();
let mut out = [0u8; 256];
let m = decode(&frame[..n], &mut out).unwrap();
assert_eq!(&out[..m], &payload[..]);
}
#[test]
fn decode_skips_a_leading_flush_end() {
let mut out = [0u8; 4];
let n = decode(&[END, b'h', b'i', END], &mut out).unwrap();
assert_eq!(&out[..n], b"hi");
}
#[test]
fn decode_tolerates_a_missing_trailing_end() {
let mut out = [0u8; 4];
let n = decode(b"hi", &mut out).unwrap();
assert_eq!(&out[..n], b"hi");
}
#[test]
fn an_invalid_escape_is_rejected() {
let mut out = [0u8; 4];
assert_eq!(
decode(&[ESC, 0x01, END], &mut out),
Err(SerialError::InvalidEscape)
);
}
#[test]
fn an_escape_at_the_end_is_truncated() {
let mut out = [0u8; 4];
assert_eq!(
decode(&[0x01, ESC], &mut out),
Err(SerialError::TruncatedFrame)
);
}
#[test]
fn encode_reports_a_full_buffer() {
let mut frame = [0u8; 2];
assert_eq!(encode(&[END], &mut frame), Err(SerialError::BufferTooSmall));
}
#[test]
fn decode_reports_a_full_buffer() {
let mut out = [0u8; 1];
assert_eq!(decode(b"hi", &mut out), Err(SerialError::BufferTooSmall));
}
#[test]
fn streaming_decoder_splits_back_to_back_frames() {
let mut decoder: SlipDecoder<16> = SlipDecoder::new();
let stream = [b'o', b'k', END, ESC, ESC_END, END];
let expected: [&[u8]; 2] = [b"ok", &[END]];
let mut seen = 0;
for &byte in &stream {
if let Some(frame) = decoder.push(byte).unwrap() {
assert_eq!(frame, expected[seen]);
seen += 1;
}
}
assert_eq!(seen, 2);
}
#[test]
fn streaming_decoder_ignores_empty_frames() {
let mut decoder: SlipDecoder<8> = SlipDecoder::new();
assert!(decoder.push(END).unwrap().is_none());
assert!(decoder.push(END).unwrap().is_none());
assert!(decoder.push(b'x').unwrap().is_none());
assert_eq!(decoder.push(END).unwrap(), Some(&b"x"[..]));
}
#[test]
fn streaming_decoder_reports_overflow_then_recovers() {
let mut decoder: SlipDecoder<2> = SlipDecoder::new();
assert!(decoder.push(b'a').unwrap().is_none());
assert!(decoder.push(b'b').unwrap().is_none());
assert_eq!(decoder.push(b'c'), Err(SerialError::BufferTooSmall));
assert!(decoder.push(b'z').unwrap().is_none());
assert_eq!(decoder.push(END).unwrap(), Some(&b"z"[..]));
}
}