#![cfg_attr(
not(any(windows, test)),
expect(
dead_code,
reason = "named-pipe framing is Windows transport plus tests"
)
)]
use std::mem::size_of;
use crate::constants::MAX_FRAME_LEN;
use crate::durability::Durability;
use crate::session_id::SessionId;
#[derive(Clone, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub(crate) enum Message {
Attach {
cols: u16,
rows: u16,
},
Input(Vec<u8>),
Resize {
cols: u16,
rows: u16,
},
Attached {
session_id: SessionId,
},
Output(Vec<u8>),
Displaced,
AppExited {
status: i32,
},
StartupOk {
session_id: SessionId,
durability: Durability,
},
StartupErr,
StartupCommit,
}
const DURABILITY_DURABLE: u8 = 1;
const DURABILITY_TIED_TO_LAUNCHER: u8 = 2;
const KIND_ATTACH: u8 = 1;
const KIND_INPUT: u8 = 2;
const KIND_RESIZE: u8 = 3;
const KIND_ATTACHED: u8 = 4;
const KIND_OUTPUT: u8 = 5;
const KIND_DISPLACED: u8 = 6;
const KIND_APP_EXITED: u8 = 7;
const KIND_STARTUP_OK: u8 = 8;
const KIND_STARTUP_ERR: u8 = 9;
const KIND_STARTUP_COMMIT: u8 = 10;
#[must_use]
pub(crate) fn encode(message: &Message) -> Vec<u8> {
let mut payload = Vec::new();
match message {
Message::Attach { cols, rows } => {
payload.push(KIND_ATTACH);
payload.extend_from_slice(&cols.to_le_bytes());
payload.extend_from_slice(&rows.to_le_bytes());
}
Message::Input(data) => {
payload.push(KIND_INPUT);
payload.extend_from_slice(data);
}
Message::Resize { cols, rows } => {
payload.push(KIND_RESIZE);
payload.extend_from_slice(&cols.to_le_bytes());
payload.extend_from_slice(&rows.to_le_bytes());
}
Message::Attached { session_id } => {
payload.push(KIND_ATTACHED);
payload.extend_from_slice(&session_id.get().to_le_bytes());
}
Message::Output(data) => {
payload.push(KIND_OUTPUT);
payload.extend_from_slice(data);
}
Message::Displaced => payload.push(KIND_DISPLACED),
Message::AppExited { status } => {
payload.push(KIND_APP_EXITED);
payload.extend_from_slice(&status.to_le_bytes());
}
Message::StartupOk {
session_id,
durability,
} => {
payload.push(KIND_STARTUP_OK);
payload.extend_from_slice(&session_id.get().to_le_bytes());
payload.push(match *durability {
Durability::Durable => DURABILITY_DURABLE,
Durability::TiedToLauncher => DURABILITY_TIED_TO_LAUNCHER,
});
}
Message::StartupErr => payload.push(KIND_STARTUP_ERR),
Message::StartupCommit => payload.push(KIND_STARTUP_COMMIT),
}
let len = u32::try_from(payload.len()).expect("frame payload fits in u32");
let mut frame = Vec::with_capacity(
size_of::<u32>()
.checked_add(payload.len())
.expect("frame length fits in usize"),
);
frame.extend_from_slice(&len.to_le_bytes());
frame.extend_from_slice(&payload);
frame
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum DecodeError {
Invalid,
}
pub(crate) fn decode_payload(payload: &[u8]) -> Result<Message, DecodeError> {
let Some((kind, rest)) = payload.split_first() else {
return Err(DecodeError::Invalid);
};
match *kind {
KIND_ATTACH => decode_size(rest).map(|(cols, rows)| Message::Attach { cols, rows }),
KIND_INPUT => Ok(Message::Input(rest.to_vec())),
KIND_RESIZE => decode_size(rest).map(|(cols, rows)| Message::Resize { cols, rows }),
KIND_ATTACHED => decode_session_id(rest).map(|session_id| Message::Attached { session_id }),
KIND_OUTPUT => Ok(Message::Output(rest.to_vec())),
KIND_DISPLACED if rest.is_empty() => Ok(Message::Displaced),
KIND_APP_EXITED => decode_i32(rest).map(|status| Message::AppExited { status }),
KIND_STARTUP_OK => {
let Some((durability, id_bytes)) = rest.split_last() else {
return Err(DecodeError::Invalid);
};
let durability = match *durability {
DURABILITY_DURABLE => Durability::Durable,
DURABILITY_TIED_TO_LAUNCHER => Durability::TiedToLauncher,
_ => return Err(DecodeError::Invalid),
};
decode_session_id(id_bytes).map(|session_id| Message::StartupOk {
session_id,
durability,
})
}
KIND_STARTUP_ERR if rest.is_empty() => Ok(Message::StartupErr),
KIND_STARTUP_COMMIT if rest.is_empty() => Ok(Message::StartupCommit),
_ => Err(DecodeError::Invalid),
}
}
#[must_use]
pub(crate) fn payload_len_ok(len: u32) -> bool {
len > 0 && len <= MAX_FRAME_LEN
}
fn decode_size(rest: &[u8]) -> Result<(u16, u16), DecodeError> {
let (cols_bytes, rest) = rest.split_at_checked(2).ok_or(DecodeError::Invalid)?;
let (rows_bytes, rest) = rest.split_at_checked(2).ok_or(DecodeError::Invalid)?;
if !rest.is_empty() {
return Err(DecodeError::Invalid);
}
let cols = u16::from_le_bytes(
cols_bytes
.try_into()
.map_err(|_error| DecodeError::Invalid)?,
);
let rows = u16::from_le_bytes(
rows_bytes
.try_into()
.map_err(|_error| DecodeError::Invalid)?,
);
Ok((cols, rows))
}
fn decode_u32(rest: &[u8]) -> Result<u32, DecodeError> {
let (bytes, rest) = rest.split_at_checked(4).ok_or(DecodeError::Invalid)?;
if !rest.is_empty() {
return Err(DecodeError::Invalid);
}
Ok(u32::from_le_bytes(
bytes.try_into().map_err(|_error| DecodeError::Invalid)?,
))
}
fn decode_session_id(rest: &[u8]) -> Result<SessionId, DecodeError> {
SessionId::from_u32(decode_u32(rest)?).ok_or(DecodeError::Invalid)
}
fn decode_i32(rest: &[u8]) -> Result<i32, DecodeError> {
Ok(i32::from_le_bytes(decode_u32(rest)?.to_le_bytes()))
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use super::*;
#[test]
fn round_trips_each_kind() {
let id = SessionId::MIN;
let messages = [
Message::Attach { cols: 80, rows: 24 },
Message::Input(b"hi".to_vec()),
Message::Resize {
cols: 120,
rows: 30,
},
Message::Attached { session_id: id },
Message::Output(b"out".to_vec()),
Message::Displaced,
Message::AppExited { status: 7 },
Message::StartupOk {
session_id: id,
durability: Durability::Durable,
},
Message::StartupOk {
session_id: id,
durability: Durability::TiedToLauncher,
},
Message::StartupErr,
Message::StartupCommit,
];
for message in messages {
let frame = encode(&message);
let (header, payload) = frame
.split_first_chunk::<4>()
.expect("encode always writes a length prefix");
let len = u32::from_le_bytes(*header);
assert!(payload_len_ok(len));
assert_eq!(payload.len(), len as usize);
assert_eq!(decode_payload(payload).unwrap(), message);
}
}
#[test]
fn rejects_empty_payload() {
assert_eq!(decode_payload(&[]).unwrap_err(), DecodeError::Invalid);
}
#[test]
fn rejects_zero_session_id() {
let mut payload = vec![KIND_ATTACHED];
payload.extend_from_slice(&0_u32.to_le_bytes());
assert_eq!(decode_payload(&payload).unwrap_err(), DecodeError::Invalid);
}
#[test]
fn payload_len_rejects_zero_and_over_cap() {
assert!(!payload_len_ok(0));
assert!(!payload_len_ok(MAX_FRAME_LEN.saturating_add(1)));
assert!(payload_len_ok(1));
assert!(payload_len_ok(64 * 1024));
assert!(payload_len_ok(MAX_FRAME_LEN));
}
#[test]
fn empty_messages_reject_trailing_bytes() {
assert_eq!(
decode_payload(&[KIND_DISPLACED, 0]).unwrap_err(),
DecodeError::Invalid
);
assert_eq!(
decode_payload(&[KIND_STARTUP_ERR, 1]).unwrap_err(),
DecodeError::Invalid
);
assert_eq!(
decode_payload(&[KIND_STARTUP_COMMIT, 1]).unwrap_err(),
DecodeError::Invalid
);
}
#[test]
fn startup_ok_rejects_a_missing_or_unknown_durability() {
let id = SessionId::MIN;
let mut without_durability = vec![KIND_STARTUP_OK];
without_durability.extend_from_slice(&id.get().to_le_bytes());
assert_eq!(
decode_payload(&without_durability).unwrap_err(),
DecodeError::Invalid
);
let mut unknown_durability = without_durability.clone();
unknown_durability.push(0);
assert_eq!(
decode_payload(&unknown_durability).unwrap_err(),
DecodeError::Invalid
);
assert_eq!(
decode_payload(&[KIND_STARTUP_OK]).unwrap_err(),
DecodeError::Invalid
);
}
#[test]
fn sized_and_numeric_payloads_reject_trailing_bytes() {
let mut attach = vec![KIND_ATTACH];
attach.extend_from_slice(&80_u16.to_le_bytes());
attach.extend_from_slice(&24_u16.to_le_bytes());
attach.push(0);
assert_eq!(decode_payload(&attach).unwrap_err(), DecodeError::Invalid);
let mut attached = vec![KIND_ATTACHED];
attached.extend_from_slice(&1_u32.to_le_bytes());
attached.push(0);
assert_eq!(decode_payload(&attached).unwrap_err(), DecodeError::Invalid);
}
}