use std::io::{self, Read, Write};
use prost::Message;
pub use running_process_protocol::broker::v1::{Frame, FrameKind, PayloadEncoding};
pub const FRAMING_VERSION_V1: u8 = 1;
pub const MAX_FRAME_SIZE_BYTES: usize = 16 * 1024 * 1024;
pub const MAX_HELLO_SIZE_BYTES: usize = 64 * 1024;
pub const ENVELOPE_VERSION: u8 = FRAMING_VERSION_V1;
pub const MAX_FRAME_BYTES: usize = MAX_FRAME_SIZE_BYTES;
pub const MAX_HELLO_BYTES: usize = MAX_HELLO_SIZE_BYTES;
pub const FRAME_HEADER_BYTES: usize = 5;
#[derive(Debug, thiserror::Error)]
pub enum FramingError {
#[error("unsupported framing version: got {got}, expected {expected}")]
UnsupportedFramingVersion {
got: u8,
expected: u8,
},
#[error("frame body too large: {body_length} bytes exceeds cap {cap}")]
FrameTooLarge {
body_length: usize,
cap: usize,
},
#[error("unexpected EOF while reading frame ({context})")]
UnexpectedEof {
context: &'static str,
},
#[error("I/O error: {0}")]
Io(#[from] io::Error),
#[error("failed to decode Frame body: {0}")]
Decode(#[from] prost::DecodeError),
}
pub fn encode_framed(frame: &Frame) -> Result<Vec<u8>, FramingError> {
let body_len = frame.encoded_len();
if body_len > MAX_FRAME_BYTES {
return Err(FramingError::FrameTooLarge {
body_length: body_len,
cap: MAX_FRAME_BYTES,
});
}
let mut wire = Vec::with_capacity(FRAME_HEADER_BYTES + body_len);
wire.push(ENVELOPE_VERSION);
wire.extend_from_slice(&(body_len as u32).to_le_bytes());
frame
.encode(&mut wire)
.expect("prost encoding into Vec cannot fail because Vec writes are infallible");
Ok(wire)
}
#[derive(Debug, Clone, PartialEq)]
pub struct DecodedFramed {
pub frame: Frame,
pub consumed: usize,
}
pub fn try_decode_framed(buf: &[u8]) -> Result<Option<DecodedFramed>, FramingError> {
if buf.is_empty() {
return Ok(None);
}
if buf[0] != ENVELOPE_VERSION {
return Err(FramingError::UnsupportedFramingVersion {
got: buf[0],
expected: ENVELOPE_VERSION,
});
}
if buf.len() < FRAME_HEADER_BYTES {
return Ok(None);
}
let body_len = u32::from_le_bytes([buf[1], buf[2], buf[3], buf[4]]) as usize;
if body_len > MAX_FRAME_BYTES {
return Err(FramingError::FrameTooLarge {
body_length: body_len,
cap: MAX_FRAME_BYTES,
});
}
let total = FRAME_HEADER_BYTES + body_len;
if buf.len() < total {
return Ok(None);
}
let frame = Frame::decode(&buf[FRAME_HEADER_BYTES..total])?;
Ok(Some(DecodedFramed {
frame,
consumed: total,
}))
}
pub fn read_frame<R: Read>(reader: &mut R) -> Result<Vec<u8>, FramingError> {
read_frame_with_cap(reader, MAX_FRAME_BYTES)
}
pub fn read_frame_with_cap<R: Read>(
reader: &mut R,
max_bytes: usize,
) -> Result<Vec<u8>, FramingError> {
let mut version_buf = [0_u8; 1];
read_exact_or_eof(reader, &mut version_buf, "framing byte")?;
let version = version_buf[0];
if version != ENVELOPE_VERSION {
return Err(FramingError::UnsupportedFramingVersion {
got: version,
expected: ENVELOPE_VERSION,
});
}
let mut len_buf = [0_u8; 4];
read_exact_or_eof(reader, &mut len_buf, "body length header")?;
let body_length = u32::from_le_bytes(len_buf) as usize;
if body_length > max_bytes {
return Err(FramingError::FrameTooLarge {
body_length,
cap: max_bytes,
});
}
let mut body = vec![0_u8; body_length];
if body_length != 0 {
read_exact_or_eof(reader, &mut body, "frame body")?;
}
Ok(body)
}
pub fn write_frame<W: Write>(writer: &mut W, body: &[u8]) -> Result<usize, FramingError> {
if body.len() > MAX_FRAME_BYTES {
return Err(FramingError::FrameTooLarge {
body_length: body.len(),
cap: MAX_FRAME_BYTES,
});
}
let body_len = body.len() as u32;
let header = [
ENVELOPE_VERSION,
(body_len & 0xFF) as u8,
((body_len >> 8) & 0xFF) as u8,
((body_len >> 16) & 0xFF) as u8,
((body_len >> 24) & 0xFF) as u8,
];
writer.write_all(&header)?;
if !body.is_empty() {
writer.write_all(body)?;
}
writer.flush()?;
Ok(header.len() + body.len())
}
fn read_exact_or_eof<R: Read>(
reader: &mut R,
buf: &mut [u8],
context: &'static str,
) -> Result<(), FramingError> {
match reader.read_exact(buf) {
Ok(()) => Ok(()),
Err(error) if error.kind() == io::ErrorKind::UnexpectedEof => {
Err(FramingError::UnexpectedEof { context })
}
Err(error) => Err(FramingError::Io(error)),
}
}
pub mod registry {
pub const PROTOCOL_VERSION: u32 = 1;
pub const CONTROL_PAYLOAD_PROTOCOL: u32 = 0x00;
pub const ADMIN_PAYLOAD_PROTOCOL: u32 = 0xAD01;
pub const BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL: u32 = 0xB232;
pub const HANDOFF_PAYLOAD_PROTOCOL: u32 = 0xD0FF;
pub const SESSION_PAYLOAD_PROTOCOL: u32 = 0x5350;
pub const CONSUMER_PAYLOAD_PROTOCOL_MIN: u32 = 0x7000;
pub const CONSUMER_PAYLOAD_PROTOCOL_MAX: u32 = 0x7EFF;
pub const PRIVATE_USE_PAYLOAD_PROTOCOL_MIN: u32 = 0xF000;
pub const PRIVATE_USE_PAYLOAD_PROTOCOL_MAX: u32 = 0xFFFF;
pub const ZCCACHE_PAYLOAD_PROTOCOL: u32 = 0x7A63;
pub const CLUD_PAYLOAD_PROTOCOL: u32 = 0x7C4C;
pub const FBUILD_PAYLOAD_PROTOCOL: u32 = 0x7EB1;
pub const FIRST_PARTY_PAYLOAD_PROTOCOLS: [u32; 4] = [
CONTROL_PAYLOAD_PROTOCOL,
ADMIN_PAYLOAD_PROTOCOL,
BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL,
HANDOFF_PAYLOAD_PROTOCOL,
];
pub const fn is_first_party(id: u32) -> bool {
let mut index = 0;
while index < FIRST_PARTY_PAYLOAD_PROTOCOLS.len() {
if FIRST_PARTY_PAYLOAD_PROTOCOLS[index] == id {
return true;
}
index += 1;
}
false
}
pub const fn is_registered_consumer_id(id: u32) -> bool {
id >= CONSUMER_PAYLOAD_PROTOCOL_MIN && id <= CONSUMER_PAYLOAD_PROTOCOL_MAX
}
pub const fn is_private_use_id(id: u32) -> bool {
id >= PRIVATE_USE_PAYLOAD_PROTOCOL_MIN && id <= PRIVATE_USE_PAYLOAD_PROTOCOL_MAX
}
}
pub use registry::{
ADMIN_PAYLOAD_PROTOCOL, BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL, CLUD_PAYLOAD_PROTOCOL,
CONTROL_PAYLOAD_PROTOCOL, FBUILD_PAYLOAD_PROTOCOL, HANDOFF_PAYLOAD_PROTOCOL, PROTOCOL_VERSION,
SESSION_PAYLOAD_PROTOCOL, ZCCACHE_PAYLOAD_PROTOCOL,
};
#[macro_export]
macro_rules! register_payload_protocol {
($(#[$meta:meta])* $vis:vis const $name:ident: u32 = $value:expr;) => {
$(#[$meta])*
$vis const $name: u32 = $value;
const _: () = {
assert!(
!$crate::frame_v1::registry::is_first_party($name),
concat!(
stringify!($name),
" collides with a first-party running-process payload protocol",
),
);
assert!(
$crate::frame_v1::registry::is_registered_consumer_id($name)
|| $crate::frame_v1::registry::is_private_use_id($name),
concat!(
stringify!($name),
" must lie in the registered-consumer range (0x7000..=0x7EFF) ",
"or the private-use range (0xF000..=0xFFFF)",
),
);
};
};
}
#[cfg(test)]
mod tests {
use super::registry::{
is_first_party, is_private_use_id, is_registered_consumer_id, ADMIN_PAYLOAD_PROTOCOL,
BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL, CLUD_PAYLOAD_PROTOCOL,
CONSUMER_PAYLOAD_PROTOCOL_MAX, CONSUMER_PAYLOAD_PROTOCOL_MIN, CONTROL_PAYLOAD_PROTOCOL,
FBUILD_PAYLOAD_PROTOCOL, HANDOFF_PAYLOAD_PROTOCOL, PRIVATE_USE_PAYLOAD_PROTOCOL_MAX,
PRIVATE_USE_PAYLOAD_PROTOCOL_MIN, PROTOCOL_VERSION, ZCCACHE_PAYLOAD_PROTOCOL,
};
use super::{
encode_framed, try_decode_framed, Frame, FrameKind, FramingError, PayloadEncoding,
ENVELOPE_VERSION, MAX_FRAME_BYTES,
};
crate::register_payload_protocol! {
const MACRO_CONSUMER_RANGE_EXAMPLE: u32 = 0x7001;
}
crate::register_payload_protocol! {
const MACRO_PRIVATE_RANGE_EXAMPLE: u32 = 0xF00D;
}
#[test]
fn payload_protocol_ids_are_pairwise_distinct() {
let registered: [(u32, &str); 4] = [
(CONTROL_PAYLOAD_PROTOCOL, "CONTROL_PAYLOAD_PROTOCOL"),
(ADMIN_PAYLOAD_PROTOCOL, "ADMIN_PAYLOAD_PROTOCOL"),
(
BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL,
"BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL",
),
(HANDOFF_PAYLOAD_PROTOCOL, "HANDOFF_PAYLOAD_PROTOCOL"),
];
for (left_index, (left_id, left_name)) in registered.iter().enumerate() {
for (right_id, right_name) in ®istered[left_index + 1..] {
assert_ne!(
left_id, right_id,
"{left_name} and {right_name} share payload-protocol id {left_id:#06X}"
);
}
}
}
#[test]
fn frozen_v1_wire_values() {
assert_eq!(PROTOCOL_VERSION, 1);
assert_eq!(CONTROL_PAYLOAD_PROTOCOL, 0x00);
assert_eq!(ADMIN_PAYLOAD_PROTOCOL, 0xAD01);
assert_eq!(BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL, 0xB232);
assert_eq!(HANDOFF_PAYLOAD_PROTOCOL, 0xD0FF);
assert_eq!(u32::from(super::FRAMING_VERSION_V1), 1);
}
#[test]
fn frozen_consumer_registry_values() {
assert_eq!(CONSUMER_PAYLOAD_PROTOCOL_MIN, 0x7000);
assert_eq!(CONSUMER_PAYLOAD_PROTOCOL_MAX, 0x7EFF);
assert_eq!(PRIVATE_USE_PAYLOAD_PROTOCOL_MIN, 0xF000);
assert_eq!(PRIVATE_USE_PAYLOAD_PROTOCOL_MAX, 0xFFFF);
assert_eq!(ZCCACHE_PAYLOAD_PROTOCOL, 0x7A63);
assert_eq!(CLUD_PAYLOAD_PROTOCOL, 0x7C4C);
assert_eq!(FBUILD_PAYLOAD_PROTOCOL, 0x7EB1);
assert!(is_first_party(BACKEND_HANDLE_PROBE_PAYLOAD_PROTOCOL));
assert!(is_registered_consumer_id(ZCCACHE_PAYLOAD_PROTOCOL));
assert!(is_private_use_id(0xF412));
}
#[test]
fn register_macro_defines_usable_constants() {
assert_eq!(MACRO_CONSUMER_RANGE_EXAMPLE, 0x7001);
assert_eq!(MACRO_PRIVATE_RANGE_EXAMPLE, 0xF00D);
}
#[test]
fn root_client_reexports_keep_frame_extensions_and_codecs() {
let frame = Frame::request(0x7A63, b"ping".to_vec()).with_request_id(42);
assert_eq!(frame.kind, FrameKind::Request as i32);
assert_eq!(frame.payload_encoding, PayloadEncoding::None as i32);
let wire = encode_framed(&frame).expect("encode");
let decoded = try_decode_framed(&wire)
.expect("decode")
.expect("complete frame");
assert_eq!(decoded.frame, frame);
assert_eq!(decoded.consumed, wire.len());
}
#[test]
fn try_decode_framed_waits_for_complete_frames() {
let wire = encode_framed(&Frame::request(0x7001, b"abc".to_vec())).expect("encode");
assert!(
try_decode_framed(&[])
.expect("empty buffer is not an error")
.is_none(),
"an empty buffer must ask for more bytes, not decode"
);
for cut in 1..wire.len() {
assert!(
try_decode_framed(&wire[..cut])
.expect("a prefix is not an error")
.is_none(),
"partial frame of {cut} of {} bytes must not decode",
wire.len()
);
}
let mut two = wire.clone();
two.extend_from_slice(&wire);
let first = try_decode_framed(&two)
.expect("decode")
.expect("first frame is complete");
assert_eq!(first.consumed, wire.len());
}
#[test]
fn try_decode_framed_rejects_foreign_version_and_oversize() {
let foreign = ENVELOPE_VERSION.wrapping_add(1);
assert!(matches!(
try_decode_framed(&[foreign, 0, 0, 0, 0]),
Err(FramingError::UnsupportedFramingVersion { got, expected })
if got == foreign && expected == ENVELOPE_VERSION
));
let mut oversize = vec![ENVELOPE_VERSION];
let claimed = u32::try_from(MAX_FRAME_BYTES).expect("cap fits u32") + 1;
oversize.extend_from_slice(&claimed.to_le_bytes());
assert!(matches!(
try_decode_framed(&oversize),
Err(FramingError::FrameTooLarge { body_length, cap })
if body_length == claimed as usize && cap == MAX_FRAME_BYTES
));
}
#[cfg(feature = "client")]
#[test]
fn root_client_reexports_keep_endpoint_extensions_and_errors() {
use crate::broker::protocol::{Endpoint, EndpointNameError};
assert!(Endpoint::unix_socket("svc", "/tmp/svc.sock").is_ok());
assert_eq!(
Endpoint::windows_pipe("svc", r"\\.\pipe\svc-pipe"),
Err(EndpointNameError::PrefixedPipeName {
got: r"\\.\pipe\svc-pipe".to_owned(),
})
);
}
}