use super::WireOp;
use crate::encoder::{SnapshotFormat, MAX_SNAPSHOT_BYTES};
use crate::pcc::QualityConfig;
use anyhow::{Context, Result};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
pub const PROTOCOL_VERSION: u8 = 5;
pub const MAX_MESSAGE_SIZE: u32 = 16 * 1024 * 1024;
pub const MAX_OPS_PER_UPDATE: u32 = 8_192;
pub const MAX_TOKEN_LEN: u16 = 64;
pub const MAX_ERROR_LEN: u16 = 1024;
pub const SNAPSHOT_CHUNK_BYTES: usize = 1024 * 1024;
pub const MAX_SNAPSHOT_CHUNKS: u32 = 128;
const K_HELLO: u8 = 0x01;
const K_E2E_OFFER: u8 = 0x12;
const K_E2E_REPLY: u8 = 0x13;
const K_REQUEST_KEYFRAME: u8 = 0x02;
const K_ACK: u8 = 0x03;
const K_SNAPSHOT_BEGIN: u8 = 0x04;
const K_SNAPSHOT_CHUNK: u8 = 0x05;
const K_PARTIAL_UPDATE: u8 = 0x06;
const K_KEEP_ALIVE: u8 = 0x07;
const K_QUALITY: u8 = 0x08;
const K_ERROR: u8 = 0x09;
const K_BYE: u8 = 0x0A;
const K_SNAPSHOT_COMMIT: u8 = 0x0B;
pub type Rev = u64;
pub type Epoch = u32;
#[derive(Debug, Clone, PartialEq)]
pub enum Message {
Hello {
token: String,
},
RequestKeyframe,
E2eOffer {
public: [u8; 32],
proof: [u8; 32],
},
E2eReply {
public: [u8; 32],
proof: [u8; 32],
},
Ack {
rev: Rev,
},
SnapshotBegin {
rev: Rev,
pts_us: u64,
epoch: Epoch,
width: u32,
height: u32,
format: SnapshotFormat,
total_len: u32,
chunks: u32,
},
SnapshotChunk {
rev: Rev,
index: u32,
data: Vec<u8>,
},
SnapshotCommit {
rev: Rev,
pts_us: u64,
epoch: Epoch,
},
PartialUpdate {
rev: Rev,
pts_us: u64,
epoch: Epoch,
ops: Vec<WireOp>,
},
KeepAlive {
rev: Rev,
},
QualityConfig(QualityConfig),
Error(String),
Bye,
}
impl Message {
pub fn encode_into(&self, out: &mut Vec<u8>) {
out.push(PROTOCOL_VERSION);
let len_at = out.len();
out.extend_from_slice(&[0u8; 4]);
self.encode_body(out);
let len = (out.len() - len_at - 4) as u32;
out[len_at..len_at + 4].copy_from_slice(&len.to_le_bytes());
}
pub fn encode(&self) -> Result<Vec<u8>> {
let mut out = Vec::with_capacity(1024);
out.push(PROTOCOL_VERSION);
let len_at = out.len();
out.extend_from_slice(&[0u8; 4]);
self.encode_body(&mut out);
let len = (out.len() - len_at - 4) as u64;
if len > MAX_MESSAGE_SIZE as u64 {
anyhow::bail!("Message too large: {len} bytes (max_message_size={MAX_MESSAGE_SIZE})");
}
let len = len as u32;
out[len_at..len_at + 4].copy_from_slice(&len.to_le_bytes());
Ok(out)
}
pub fn encoded_len(&self) -> Result<usize> {
let mut buf = Vec::new();
self.encode_body(&mut buf);
let len = buf.len();
drop(buf);
if len as u64 > MAX_MESSAGE_SIZE as u64 {
anyhow::bail!("Message too large: {len} bytes (max_message_size={MAX_MESSAGE_SIZE})");
}
Ok(len + 5)
}
fn encode_body(&self, out: &mut Vec<u8>) {
match self {
Message::Hello { token } => {
out.push(K_HELLO);
let bytes = token.as_bytes();
out.extend_from_slice(&(bytes.len() as u16).to_le_bytes());
out.extend_from_slice(bytes);
}
Message::RequestKeyframe => out.push(K_REQUEST_KEYFRAME),
Message::E2eOffer { public, proof } => {
out.push(K_E2E_OFFER);
out.extend_from_slice(public);
out.extend_from_slice(proof);
}
Message::E2eReply { public, proof } => {
out.push(K_E2E_REPLY);
out.extend_from_slice(public);
out.extend_from_slice(proof);
}
Message::Ack { rev } => {
out.push(K_ACK);
out.extend_from_slice(&rev.to_le_bytes());
}
Message::SnapshotBegin {
rev,
pts_us,
epoch,
width,
height,
format,
total_len,
chunks,
} => {
out.push(K_SNAPSHOT_BEGIN);
out.extend_from_slice(&rev.to_le_bytes());
out.extend_from_slice(&pts_us.to_le_bytes());
out.extend_from_slice(&epoch.to_le_bytes());
out.extend_from_slice(&width.to_le_bytes());
out.extend_from_slice(&height.to_le_bytes());
out.push(u8::from(*format));
out.extend_from_slice(&total_len.to_le_bytes());
out.extend_from_slice(&chunks.to_le_bytes());
}
Message::SnapshotChunk { rev, index, data } => {
out.push(K_SNAPSHOT_CHUNK);
out.extend_from_slice(&rev.to_le_bytes());
out.extend_from_slice(&index.to_le_bytes());
out.extend_from_slice(&(data.len() as u32).to_le_bytes());
out.extend_from_slice(data);
}
Message::SnapshotCommit { rev, pts_us, epoch } => {
out.push(K_SNAPSHOT_COMMIT);
out.extend_from_slice(&rev.to_le_bytes());
out.extend_from_slice(&pts_us.to_le_bytes());
out.extend_from_slice(&epoch.to_le_bytes());
}
Message::PartialUpdate {
rev,
pts_us,
epoch,
ops,
} => {
out.push(K_PARTIAL_UPDATE);
out.extend_from_slice(&rev.to_le_bytes());
out.extend_from_slice(&pts_us.to_le_bytes());
out.extend_from_slice(&epoch.to_le_bytes());
out.extend_from_slice(&(ops.len() as u32).to_le_bytes());
for op in ops {
op.encode_into(out);
}
}
Message::KeepAlive { rev } => {
out.push(K_KEEP_ALIVE);
out.extend_from_slice(&rev.to_le_bytes());
}
Message::QualityConfig(cfg) => {
out.push(K_QUALITY);
out.extend_from_slice(&cfg.target_fps.to_le_bytes());
out.extend_from_slice(&cfg.max_fps.to_le_bytes());
out.extend_from_slice(&cfg.quality.to_le_bytes());
}
Message::Error(text) => {
out.push(K_ERROR);
let bytes = text.as_bytes();
let n = bytes.len().min(MAX_ERROR_LEN as usize);
out.extend_from_slice(&(n as u16).to_le_bytes());
out.extend_from_slice(&bytes[..n]);
}
Message::Bye => out.push(K_BYE),
}
}
pub fn decode(bytes: &[u8]) -> Result<Self> {
if bytes.len() < 5 {
anyhow::bail!("Message too short: {} bytes (need >= 5)", bytes.len());
}
let version = bytes[0];
if version != PROTOCOL_VERSION {
anyhow::bail!("Protocol version mismatch: expected {PROTOCOL_VERSION}, got {version}");
}
let len = u32::from_le_bytes(bytes[1..5].try_into().unwrap());
if len > MAX_MESSAGE_SIZE {
anyhow::bail!(
"Framed message too large: {len} bytes (max_message_size={MAX_MESSAGE_SIZE})"
);
}
let payload = &bytes[5..];
if payload.len() != len as usize {
anyhow::bail!(
"Message length mismatch: header says {len}, got {}",
payload.len()
);
}
let mut cur = Cursor::new(payload);
let msg = cur.message()?;
if !cur.is_empty() {
anyhow::bail!("{} trailing bytes after message body", cur.remaining());
}
Ok(msg)
}
pub fn rev(&self) -> Option<Rev> {
match self {
Message::SnapshotCommit { rev, .. }
| Message::PartialUpdate { rev, .. }
| Message::KeepAlive { rev }
| Message::Ack { rev } => Some(*rev),
_ => None,
}
}
pub async fn write_framed<W: AsyncWrite + Unpin>(&self, writer: &mut W) -> Result<()> {
let encoded = self.encode()?;
writer
.write_all(&(encoded.len() as u32).to_le_bytes())
.await?;
writer.write_all(&encoded).await?;
writer.flush().await?;
Ok(())
}
pub async fn read_envelope<R: AsyncRead + Unpin>(reader: &mut R) -> Result<Vec<u8>> {
let mut len_buf = [0u8; 4];
reader.read_exact(&mut len_buf).await?;
let len = read_len_prefix(&len_buf)?;
let mut payload = vec![0u8; len];
reader.read_exact(&mut payload).await?;
Ok(payload)
}
pub async fn read_framed<R: AsyncRead + Unpin>(reader: &mut R) -> Result<Self> {
Message::decode(&Message::read_envelope(reader).await?)
}
}
pub fn peek_rev(envelope: &[u8]) -> Option<Rev> {
if envelope.len() < 5 || envelope[0] != PROTOCOL_VERSION {
return None;
}
let body = &envelope[5..];
match *body.first()? {
K_SNAPSHOT_BEGIN | K_SNAPSHOT_CHUNK | K_PARTIAL_UPDATE | K_KEEP_ALIVE
| K_SNAPSHOT_COMMIT => {
if body.len() < 9 {
return None;
}
Some(u64::from_le_bytes(body[1..9].try_into().ok()?))
}
_ => None,
}
}
pub fn read_len_prefix(len_buf: &[u8; 4]) -> Result<usize> {
let len = u32::from_le_bytes(*len_buf);
if len > MAX_MESSAGE_SIZE + 5 {
anyhow::bail!(
"Framed message too large: {len} bytes (max_message_size={MAX_MESSAGE_SIZE})"
);
}
Ok(len as usize)
}
pub struct Cursor<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Cursor<'a> {
pub fn new(body: &'a [u8]) -> Self {
Self { buf: body, pos: 0 }
}
pub fn remaining(&self) -> usize {
self.buf.len() - self.pos
}
pub fn is_empty(&self) -> bool {
self.remaining() == 0
}
pub fn take(&mut self, n: usize) -> Result<&'a [u8]> {
if n > self.remaining() {
anyhow::bail!(
"Truncated message: need {n} more bytes, have {}",
self.remaining()
);
}
let out = &self.buf[self.pos..self.pos + n];
self.pos += n;
Ok(out)
}
pub fn u8(&mut self) -> Result<u8> {
Ok(self.take(1)?[0])
}
fn u16(&mut self) -> Result<u16> {
Ok(u16::from_le_bytes(self.take(2)?.try_into().unwrap()))
}
fn u32(&mut self) -> Result<u32> {
Ok(u32::from_le_bytes(self.take(4)?.try_into().unwrap()))
}
fn u64(&mut self) -> Result<u64> {
Ok(u64::from_le_bytes(self.take(8)?.try_into().unwrap()))
}
fn f32(&mut self) -> Result<f32> {
Ok(f32::from_le_bytes(self.take(4)?.try_into().unwrap()))
}
fn alloc_len(&mut self, max: u64, what: &str) -> Result<usize> {
let n = self.u32()? as u64;
if n > max {
anyhow::bail!("{what} too large: {n} bytes (max={max})");
}
Ok(n as usize)
}
fn message(&mut self) -> Result<Message> {
let kind = self.u8()?;
let msg = match kind {
K_HELLO => {
let n = self.u16()? as usize;
if n > MAX_TOKEN_LEN as usize {
anyhow::bail!("token too large: {n} bytes (max={MAX_TOKEN_LEN})");
}
let token = std::str::from_utf8(self.take(n)?)
.context("token is not valid UTF-8")?
.to_string();
Message::Hello { token }
}
K_REQUEST_KEYFRAME => Message::RequestKeyframe,
K_E2E_OFFER => {
let mut public = [0u8; 32];
let mut proof = [0u8; 32];
public.copy_from_slice(self.take(32)?);
proof.copy_from_slice(self.take(32)?);
Message::E2eOffer { public, proof }
}
K_E2E_REPLY => {
let mut public = [0u8; 32];
let mut proof = [0u8; 32];
public.copy_from_slice(self.take(32)?);
proof.copy_from_slice(self.take(32)?);
Message::E2eReply { public, proof }
}
K_ACK => Message::Ack { rev: self.u64()? },
K_SNAPSHOT_BEGIN => {
let rev = self.u64()?;
let pts_us = self.u64()?;
let epoch = self.u32()?;
let width = self.u32()?;
let height = self.u32()?;
crate::pcc::types::rgb_len(width, height).map_err(|e| {
anyhow::anyhow!("invalid snapshot geometry {width}x{height}: {e}")
})?;
let format = SnapshotFormat::from_u8(self.u8()?)?;
let total_len = self.alloc_len(MAX_SNAPSHOT_BYTES as u64, "snapshot")? as u32;
let chunks = self.u32()?;
if chunks == 0 || chunks > MAX_SNAPSHOT_CHUNKS {
anyhow::bail!(
"snapshot chunk count {chunks} outside 1..={MAX_SNAPSHOT_CHUNKS}"
);
}
Message::SnapshotBegin {
rev,
pts_us,
epoch,
width,
height,
format,
total_len,
chunks,
}
}
K_SNAPSHOT_CHUNK => {
let rev = self.u64()?;
let index = self.u32()?;
let n = self.alloc_len(SNAPSHOT_CHUNK_BYTES as u64, "snapshot chunk")?;
Message::SnapshotChunk {
rev,
index,
data: self.take(n)?.to_vec(),
}
}
K_PARTIAL_UPDATE => {
let rev = self.u64()?;
let pts_us = self.u64()?;
let epoch = self.u32()?;
let count = self.u32()?;
if count > MAX_OPS_PER_UPDATE {
anyhow::bail!(
"too many ops in one update: {count} (max_ops_per_update={MAX_OPS_PER_UPDATE})"
);
}
let mut ops = Vec::with_capacity((count as usize).min(1024));
for _ in 0..count {
ops.push(self.op()?);
}
Message::PartialUpdate {
rev,
pts_us,
epoch,
ops,
}
}
K_KEEP_ALIVE => Message::KeepAlive { rev: self.u64()? },
K_QUALITY => Message::QualityConfig(QualityConfig {
target_fps: self.u32()?,
max_fps: self.u32()?,
quality: self.f32()?,
}),
K_ERROR => {
let n = self.u16()?;
if n > MAX_ERROR_LEN {
anyhow::bail!("Error message too large: {n} bytes (max={MAX_ERROR_LEN})");
}
let text = std::str::from_utf8(self.take(n as usize)?)
.context("Error message is not valid UTF-8")?
.to_string();
Message::Error(text)
}
K_SNAPSHOT_COMMIT => Message::SnapshotCommit {
rev: self.u64()?,
pts_us: self.u64()?,
epoch: self.u32()?,
},
K_BYE => Message::Bye,
other => anyhow::bail!("Unknown message kind: 0x{other:02x}"),
};
Ok(msg)
}
fn op(&mut self) -> Result<WireOp> {
let x = self.u32()?;
let y = self.u32()?;
let width = self.u32()?;
let height = self.u32()?;
let kind = self.u8()?;
let op = match kind {
crate::network::OP_RECT => {
let expected = (width as u64) * (height as u64) * 3;
if expected > crate::pcc::types::MAX_FRAME_BYTES as u64 {
anyhow::bail!(
"rect too large: {width}x{height} = {expected} raw bytes \
(max_frame_bytes={})",
crate::pcc::types::MAX_FRAME_BYTES
);
}
let n = self.alloc_len(expected + 16, "rect payload")?;
let compressed = self.take(n)?.to_vec();
WireOp::Rect {
x,
y,
width,
height,
compressed,
}
}
crate::network::OP_FILL => {
let n = self.alloc_len(3, "fill colour")?;
if n != 3 {
anyhow::bail!("fill colour must be 3 bytes, got {n}");
}
let b = self.take(3)?;
WireOp::Fill {
x,
y,
width,
height,
color: [b[0], b[1], b[2]],
}
}
crate::network::OP_COPY => {
let src_x = self.u32()?;
let src_y = self.u32()?;
WireOp::Copy {
x,
y,
width,
height,
src_x,
src_y,
}
}
other => anyhow::bail!("Unknown op kind: 0x{other:02x}"),
};
Ok(op)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn roundtrip(msg: &Message) -> Message {
let bytes = msg.encode().unwrap();
Message::decode(&bytes).unwrap()
}
#[test]
fn every_message_roundtrips() {
let msgs = vec![
Message::Hello {
token: "ABC123".into(),
},
Message::RequestKeyframe,
Message::Ack { rev: 42 },
Message::SnapshotBegin {
rev: 7,
pts_us: 0,
epoch: 1,
width: 2,
height: 2,
format: SnapshotFormat::Raw,
total_len: 12,
chunks: 1,
},
Message::SnapshotChunk {
rev: 7,
index: 0,
data: vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12],
},
Message::SnapshotCommit {
rev: 7,
pts_us: 0,
epoch: 1,
},
Message::PartialUpdate {
rev: 8,
pts_us: 0,
epoch: 1,
ops: vec![
WireOp::Rect {
x: 1,
y: 1,
width: 1,
height: 1,
compressed: vec![7, 7, 7],
},
WireOp::Fill {
x: 0,
y: 0,
width: 4,
height: 4,
color: [9, 8, 7],
},
WireOp::Copy {
x: 2,
y: 2,
width: 4,
height: 4,
src_x: 40,
src_y: 40,
},
],
},
Message::KeepAlive { rev: 99 },
Message::QualityConfig(QualityConfig::default()),
Message::Error("boom".into()),
Message::Bye,
];
for msg in msgs {
assert_eq!(roundtrip(&msg), msg, "roundtrip failed for {msg:?}");
}
}
#[test]
fn version_mismatch_is_rejected() {
let mut bytes = Message::Bye.encode().unwrap();
bytes[0] = 2;
let err = Message::decode(&bytes).unwrap_err().to_string();
assert!(err.contains("version mismatch"), "unhelpful: {err}");
}
#[test]
fn hostile_length_prefix_never_allocates() {
let err = read_len_prefix(&[0xFF, 0xFF, 0xFF, 0xFF])
.unwrap_err()
.to_string();
assert!(err.contains("max_message_size"), "unhelpful: {err}");
}
#[test]
fn truncated_body_is_rejected() {
let bytes = Message::Ack { rev: 5 }.encode().unwrap();
let err = Message::decode(&bytes[..bytes.len() - 3])
.unwrap_err()
.to_string();
assert!(
err.contains("Truncated") || err.contains("mismatch"),
"got: {err}"
);
}
#[test]
fn unknown_kind_is_rejected() {
let mut bytes = Message::Bye.encode().unwrap();
bytes[5] = 0xEE;
assert!(Message::decode(&bytes).is_err());
}
#[test]
fn snapshot_begin_with_impossible_geometry_is_rejected() {
let msg = Message::SnapshotBegin {
rev: 1,
pts_us: 0,
epoch: 0,
width: 60_000,
height: 60_000,
format: SnapshotFormat::Png,
total_len: 10,
chunks: 1,
};
let bytes = msg.encode().unwrap();
let err = Message::decode(&bytes).unwrap_err().to_string();
assert!(err.contains("max_frame_bytes"), "unhelpful: {err}");
}
#[test]
fn snapshot_begin_with_zero_chunks_is_rejected() {
let msg = Message::SnapshotBegin {
rev: 1,
pts_us: 0,
epoch: 0,
width: 2,
height: 2,
format: SnapshotFormat::Raw,
total_len: 12,
chunks: 0,
};
let bytes = msg.encode().unwrap();
assert!(Message::decode(&bytes).is_err());
}
#[test]
fn an_oversized_chunk_is_rejected_before_allocation() {
let mut payload = vec![K_SNAPSHOT_CHUNK];
payload.extend_from_slice(&1u64.to_le_bytes());
payload.extend_from_slice(&0u32.to_le_bytes());
payload.extend_from_slice(&(SNAPSHOT_CHUNK_BYTES as u32 + 1).to_le_bytes());
let mut bytes = vec![PROTOCOL_VERSION];
bytes.extend_from_slice(&(payload.len() as u32).to_le_bytes());
bytes.extend_from_slice(&payload);
let err = Message::decode(&bytes).unwrap_err().to_string();
assert!(err.contains("snapshot chunk"), "unhelpful: {err}");
}
#[test]
fn oversized_op_count_is_rejected() {
let mut payload = vec![K_PARTIAL_UPDATE];
payload.extend_from_slice(&1u64.to_le_bytes()); payload.extend_from_slice(&0u64.to_le_bytes()); payload.extend_from_slice(&0u32.to_le_bytes()); payload.extend_from_slice(&(MAX_OPS_PER_UPDATE + 1).to_le_bytes());
let mut bytes = vec![PROTOCOL_VERSION];
bytes.extend_from_slice(&(payload.len() as u32).to_le_bytes());
bytes.extend_from_slice(&payload);
let err = Message::decode(&bytes).unwrap_err().to_string();
assert!(err.contains("max_ops_per_update"), "unhelpful: {err}");
}
}