use super::error::HandshakeError;
pub(crate) struct SendBuffer<'a> {
data: &'a mut [u8],
cursor: usize,
}
impl<'a> SendBuffer<'a> {
pub(crate) fn new(data: &'a mut [u8]) -> Self {
Self { data, cursor: 0 }
}
pub(crate) fn write(&mut self, bytes: &[u8]) {
let end = self.cursor + bytes.len();
assert!(
end <= self.data.len(),
"SendBuffer overflow: need {} bytes at offset {}, buffer is {} bytes",
bytes.len(),
self.cursor,
self.data.len(),
);
self.data[self.cursor..end].copy_from_slice(bytes);
self.cursor = end;
}
pub(crate) fn reserve(&mut self, len: usize) -> &mut [u8] {
let end = self.cursor + len;
assert!(
end <= self.data.len(),
"SendBuffer overflow: need {} bytes at offset {}, buffer is {} bytes",
len,
self.cursor,
self.data.len(),
);
let slice = &mut self.data[self.cursor..end];
self.cursor = end;
slice
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.cursor
}
#[cfg(test)]
pub(crate) fn written(&self) -> &[u8] {
&self.data[..self.cursor]
}
pub(crate) fn finish(self) -> &'a [u8] {
let len = self.cursor;
&self.data[..len]
}
}
pub(crate) struct RecvBuffer<'a> {
data: &'a [u8],
cursor: usize,
}
impl<'a> RecvBuffer<'a> {
pub(crate) fn new(data: &'a [u8]) -> Self {
Self { data, cursor: 0 }
}
pub(crate) fn read(&mut self, len: usize) -> Result<&'a [u8], HandshakeError> {
let end = self.cursor + len;
if end > self.data.len() {
return Err(HandshakeError::MessageTooShort);
}
let slice = &self.data[self.cursor..end];
self.cursor = end;
Ok(slice)
}
pub(crate) fn remaining(&self) -> &'a [u8] {
&self.data[self.cursor..]
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn send_write_fills_buffer() {
let mut buf = [0u8; 10];
let mut sb = SendBuffer::new(&mut buf);
sb.write(b"hello");
sb.write(b"world");
assert_eq!(sb.len(), 10);
assert_eq!(sb.written(), b"helloworld");
}
#[test]
fn send_reserve_returns_mutable_slice() {
let mut buf = [0u8; 6];
let mut sb = SendBuffer::new(&mut buf);
sb.write(b"AB");
let slot = sb.reserve(4);
slot.copy_from_slice(b"CDEF");
assert_eq!(sb.written(), b"ABCDEF");
}
#[test]
#[should_panic(expected = "SendBuffer overflow")]
fn send_write_overflow_panics() {
let mut buf = [0u8; 3];
let mut sb = SendBuffer::new(&mut buf);
sb.write(b"toolong");
}
#[test]
#[should_panic(expected = "SendBuffer overflow")]
fn send_reserve_overflow_panics() {
let mut buf = [0u8; 3];
let mut sb = SendBuffer::new(&mut buf);
sb.reserve(4);
}
#[test]
fn recv_read_advances_cursor() {
let data = b"helloworld";
let mut rb = RecvBuffer::new(data);
let first = rb.read(5).unwrap();
assert_eq!(first, b"hello");
let second = rb.read(5).unwrap();
assert_eq!(second, b"world");
assert!(rb.remaining().is_empty());
}
#[test]
fn recv_read_too_much_returns_error() {
let data = b"short";
let mut rb = RecvBuffer::new(data);
let err = rb.read(10).unwrap_err();
assert!(matches!(err, HandshakeError::MessageTooShort));
}
#[test]
fn recv_remaining_returns_tail() {
let data = b"abcdef";
let mut rb = RecvBuffer::new(data);
let _ = rb.read(2).unwrap();
assert_eq!(rb.remaining(), b"cdef");
}
#[test]
fn send_zero_size_buffer() {
let mut buf = [0u8; 0];
let sb = SendBuffer::new(&mut buf);
assert_eq!(sb.len(), 0);
assert_eq!(sb.written(), b"");
}
#[test]
fn recv_empty_buffer() {
let data = b"";
let rb = RecvBuffer::new(data);
assert!(rb.remaining().is_empty());
}
}