use std::io::{self, Read, Write};
pub const KIND_CONTROL: u8 = 1;
pub const KIND_SCREEN: u8 = 2;
pub const KIND_HELLO: u8 = 3;
const MAX_FRAME: u32 = 64 * 1024 * 1024;
pub fn write_frame(w: &mut impl Write, kind: u8, payload: &[u8]) -> io::Result<()> {
let len = u32::try_from(payload.len())
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "frame too large"))?;
w.write_all(&len.to_be_bytes())?;
w.write_all(&[kind])?;
w.write_all(payload)?;
w.flush()
}
pub fn read_frame(r: &mut impl Read) -> io::Result<(u8, Vec<u8>)> {
let mut len_buf = [0u8; 4];
r.read_exact(&mut len_buf)?;
let len = u32::from_be_bytes(len_buf);
if len > MAX_FRAME {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"frame too large",
));
}
let mut kind = [0u8; 1];
r.read_exact(&mut kind)?;
let mut payload = vec![0u8; len as usize];
r.read_exact(&mut payload)?;
Ok((kind[0], payload))
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn frames_round_trip_back_to_back() {
let mut buf: Vec<u8> = Vec::new();
write_frame(&mut buf, KIND_CONTROL, b"hello").unwrap();
write_frame(&mut buf, KIND_SCREEN, b"").unwrap();
write_frame(&mut buf, KIND_CONTROL, &[0u8, 255, 7, 128]).unwrap();
let mut cur = Cursor::new(buf);
assert_eq!(
read_frame(&mut cur).unwrap(),
(KIND_CONTROL, b"hello".to_vec())
);
assert_eq!(read_frame(&mut cur).unwrap(), (KIND_SCREEN, Vec::new()));
assert_eq!(
read_frame(&mut cur).unwrap(),
(KIND_CONTROL, vec![0, 255, 7, 128])
);
assert!(read_frame(&mut cur).is_err());
}
#[test]
fn split_read_reassembles() {
let mut whole: Vec<u8> = Vec::new();
write_frame(&mut whole, KIND_CONTROL, b"split me").unwrap();
struct Trickle<'a>(&'a [u8], usize);
impl Read for Trickle<'_> {
fn read(&mut self, out: &mut [u8]) -> io::Result<usize> {
if self.1 >= self.0.len() || out.is_empty() {
return Ok(0);
}
out[0] = self.0[self.1];
self.1 += 1;
Ok(1)
}
}
let mut t = Trickle(&whole, 0);
assert_eq!(
read_frame(&mut t).unwrap(),
(KIND_CONTROL, b"split me".to_vec())
);
}
}