use std::{
collections::{HashMap, HashSet, VecDeque},
time::Duration,
};
pub const CARRIER_MAGIC: [u8; 4] = *b"ORIC";
pub const CARRIER_VERSION: u8 = 1;
pub const CARRIER_HEADER_LEN: usize = 51;
pub const MAX_INNER_PACKET_BYTES: usize = 1_200;
pub const DEFAULT_PACKET_QUEUE_CAPACITY: usize = 256;
pub const MAX_SEGMENT_COUNT: u16 = 64;
pub const DEFAULT_REASSEMBLY_TIMEOUT: Duration = Duration::from_secs(5);
const RECENT_PACKET_IDS: usize = DEFAULT_PACKET_QUEUE_CAPACITY * 2;
pub const CARRIER_CONTROL_PACKET_ID: u64 = u64::MAX;
const CARRIER_CONTROL_MAGIC: &[u8; 8] = b"ORICTRL1";
const CARRIER_CONTROL_TAG_BYTES: usize = 32;
const CARRIER_CONTROL_PAYLOAD_BYTES: usize =
CARRIER_CONTROL_MAGIC.len() + 1 + CARRIER_CONTROL_TAG_BYTES;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CarrierControl {
SessionTokenRevoked,
SessionTokenRevokedAck,
}
impl CarrierControl {
pub fn lifecycle_reason(self) -> Option<&'static str> {
match self {
Self::SessionTokenRevoked => {
Some(crate::lifecycle_reason::REASON_SESSION_TOKEN_REVOKED)
}
Self::SessionTokenRevokedAck => None,
}
}
fn code(self) -> u8 {
match self {
Self::SessionTokenRevoked => 1,
Self::SessionTokenRevokedAck => 2,
}
}
fn from_code(code: u8) -> Option<Self> {
match code {
1 => Some(Self::SessionTokenRevoked),
2 => Some(Self::SessionTokenRevokedAck),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CarrierFrameHeader {
pub magic: [u8; 4],
pub version: u8,
pub transport_id: u64,
pub carrier_session_id: [u8; 16],
pub transport_generation: u64,
pub packet_id: u64,
pub segment_index: u16,
pub segment_count: u16,
pub payload_len: u16,
}
impl CarrierFrameHeader {
pub fn encode(self) -> [u8; CARRIER_HEADER_LEN] {
let mut output = [0_u8; CARRIER_HEADER_LEN];
output[0..4].copy_from_slice(&self.magic);
output[4] = self.version;
output[5..13].copy_from_slice(&self.transport_id.to_be_bytes());
output[13..29].copy_from_slice(&self.carrier_session_id);
output[29..37].copy_from_slice(&self.transport_generation.to_be_bytes());
output[37..45].copy_from_slice(&self.packet_id.to_be_bytes());
output[45..47].copy_from_slice(&self.segment_index.to_be_bytes());
output[47..49].copy_from_slice(&self.segment_count.to_be_bytes());
output[49..51].copy_from_slice(&self.payload_len.to_be_bytes());
output
}
pub fn decode(input: &[u8]) -> Result<Self, CarrierFrameError> {
if input.len() < CARRIER_HEADER_LEN {
return Err(CarrierFrameError::TruncatedHeader {
actual: input.len(),
});
}
Ok(Self {
magic: input[0..4].try_into().expect("fixed magic slice"),
version: input[4],
transport_id: u64::from_be_bytes(
input[5..13].try_into().expect("fixed transport id slice"),
),
carrier_session_id: input[13..29].try_into().expect("fixed session slice"),
transport_generation: u64::from_be_bytes(
input[29..37].try_into().expect("fixed generation slice"),
),
packet_id: u64::from_be_bytes(input[37..45].try_into().expect("fixed packet id slice")),
segment_index: u16::from_be_bytes(
input[45..47].try_into().expect("fixed segment index slice"),
),
segment_count: u16::from_be_bytes(
input[47..49].try_into().expect("fixed segment count slice"),
),
payload_len: u16::from_be_bytes(
input[49..51]
.try_into()
.expect("fixed payload length slice"),
),
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CarrierFrame {
pub header: CarrierFrameHeader,
pub payload: Vec<u8>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CarrierFrameExpectation {
pub transport_id: u64,
pub carrier_session_id: [u8; 16],
pub transport_generation: u64,
}
impl CarrierFrame {
pub fn encode(&self) -> Result<Vec<u8>, CarrierFrameError> {
validate_segment_metadata(&self.header)?;
if self.payload.len() != usize::from(self.header.payload_len) {
return Err(CarrierFrameError::PayloadLength {
declared: usize::from(self.header.payload_len),
actual: self.payload.len(),
});
}
let mut output = Vec::with_capacity(CARRIER_HEADER_LEN + self.payload.len());
output.extend_from_slice(&self.header.encode());
output.extend_from_slice(&self.payload);
Ok(output)
}
pub fn decode(
input: &[u8],
expected: CarrierFrameExpectation,
) -> Result<Self, CarrierFrameError> {
let header = CarrierFrameHeader::decode(input)?;
if header.magic != CARRIER_MAGIC {
return Err(CarrierFrameError::WrongMagic);
}
if header.version != CARRIER_VERSION {
return Err(CarrierFrameError::UnsupportedVersion(header.version));
}
if header.transport_id != expected.transport_id {
return Err(CarrierFrameError::WrongTransport {
expected: expected.transport_id,
actual: header.transport_id,
});
}
if header.carrier_session_id != expected.carrier_session_id {
return Err(CarrierFrameError::WrongSession);
}
if header.transport_generation != expected.transport_generation {
return Err(CarrierFrameError::StaleGeneration {
expected: expected.transport_generation,
actual: header.transport_generation,
});
}
validate_segment_metadata(&header)?;
let payload = &input[CARRIER_HEADER_LEN..];
if payload.len() != usize::from(header.payload_len) {
return Err(CarrierFrameError::PayloadLength {
declared: usize::from(header.payload_len),
actual: payload.len(),
});
}
Ok(Self {
header,
payload: payload.to_vec(),
})
}
pub fn terminal_control(
expected: CarrierFrameExpectation,
key: &[u8; 32],
control: CarrierControl,
) -> Self {
let mut payload = Vec::with_capacity(CARRIER_CONTROL_PAYLOAD_BYTES);
payload.extend_from_slice(CARRIER_CONTROL_MAGIC);
payload.push(control.code());
payload.extend_from_slice(carrier_control_tag(expected, key, control).as_bytes());
Self {
header: CarrierFrameHeader {
magic: CARRIER_MAGIC,
version: CARRIER_VERSION,
transport_id: expected.transport_id,
carrier_session_id: expected.carrier_session_id,
transport_generation: expected.transport_generation,
packet_id: CARRIER_CONTROL_PACKET_ID,
segment_index: 0,
segment_count: 1,
payload_len: u16::try_from(payload.len()).expect("fixed carrier control length"),
},
payload,
}
}
pub fn terminal_control_kind(
&self,
expected: CarrierFrameExpectation,
key: &[u8; 32],
) -> Result<Option<CarrierControl>, CarrierFrameError> {
if self.header.packet_id != CARRIER_CONTROL_PACKET_ID {
return Ok(None);
}
if self.payload.len() != CARRIER_CONTROL_PAYLOAD_BYTES
|| &self.payload[..CARRIER_CONTROL_MAGIC.len()] != CARRIER_CONTROL_MAGIC
{
return Err(CarrierFrameError::InvalidControl);
}
let control = CarrierControl::from_code(self.payload[CARRIER_CONTROL_MAGIC.len()])
.ok_or(CarrierFrameError::InvalidControl)?;
let expected_tag = carrier_control_tag(expected, key, control);
let actual_tag = &self.payload[CARRIER_CONTROL_MAGIC.len() + 1..];
let mismatch = expected_tag
.as_bytes()
.iter()
.zip(actual_tag)
.fold(0_u8, |difference, (expected, actual)| {
difference | (expected ^ actual)
});
if mismatch != 0 {
return Err(CarrierFrameError::InvalidControlAuthentication);
}
Ok(Some(control))
}
}
fn carrier_control_tag(
expected: CarrierFrameExpectation,
key: &[u8; 32],
control: CarrierControl,
) -> blake3::Hash {
let mut input = Vec::with_capacity(64);
input.extend_from_slice(b"openrtc/iroh-carrier/control/v1");
input.extend_from_slice(&expected.transport_id.to_be_bytes());
input.extend_from_slice(&expected.carrier_session_id);
input.extend_from_slice(&expected.transport_generation.to_be_bytes());
input.push(control.code());
blake3::keyed_hash(key, &input)
}
fn validate_segment_metadata(header: &CarrierFrameHeader) -> Result<(), CarrierFrameError> {
if header.segment_count == 0 || header.segment_count > MAX_SEGMENT_COUNT {
return Err(CarrierFrameError::InvalidSegmentCount(header.segment_count));
}
if header.segment_index >= header.segment_count {
return Err(CarrierFrameError::InvalidSegmentIndex {
index: header.segment_index,
count: header.segment_count,
});
}
Ok(())
}
pub fn segment_packet(
packet: &[u8],
expected: CarrierFrameExpectation,
packet_id: u64,
carrier_message_ceiling: usize,
) -> Result<Vec<CarrierFrame>, CarrierFrameError> {
if packet_id == CARRIER_CONTROL_PACKET_ID {
return Err(CarrierFrameError::ReservedPacketId);
}
if packet.is_empty() || packet.len() > MAX_INNER_PACKET_BYTES {
return Err(CarrierFrameError::PacketSize(packet.len()));
}
let segment_payload_ceiling = carrier_message_ceiling
.checked_sub(CARRIER_HEADER_LEN)
.filter(|ceiling| *ceiling > 0)
.ok_or(CarrierFrameError::CarrierMtu(carrier_message_ceiling))?;
if segment_payload_ceiling > usize::from(u16::MAX) {
return Err(CarrierFrameError::CarrierMtu(carrier_message_ceiling));
}
let segment_count = packet.len().div_ceil(segment_payload_ceiling);
let segment_count =
u16::try_from(segment_count).map_err(|_| CarrierFrameError::TooManySegments)?;
if segment_count > MAX_SEGMENT_COUNT {
return Err(CarrierFrameError::TooManySegments);
}
packet
.chunks(segment_payload_ceiling)
.enumerate()
.map(|(segment_index, payload)| {
Ok(CarrierFrame {
header: CarrierFrameHeader {
magic: CARRIER_MAGIC,
version: CARRIER_VERSION,
transport_id: expected.transport_id,
carrier_session_id: expected.carrier_session_id,
transport_generation: expected.transport_generation,
packet_id,
segment_index: u16::try_from(segment_index)
.map_err(|_| CarrierFrameError::TooManySegments)?,
segment_count,
payload_len: u16::try_from(payload.len())
.map_err(|_| CarrierFrameError::CarrierMtu(carrier_message_ceiling))?,
},
payload: payload.to_vec(),
})
})
.collect()
}
#[derive(Debug)]
struct PartialPacket {
created_at_ms: u64,
segments: Vec<Option<Vec<u8>>>,
total_bytes: usize,
}
#[derive(Debug)]
pub struct CarrierReassembler {
partial: HashMap<u64, PartialPacket>,
recent: HashSet<u64>,
recent_order: VecDeque<u64>,
max_in_flight_packets: usize,
timeout_ms: u64,
}
impl Default for CarrierReassembler {
fn default() -> Self {
Self::new(DEFAULT_PACKET_QUEUE_CAPACITY, DEFAULT_REASSEMBLY_TIMEOUT)
}
}
impl CarrierReassembler {
pub fn new(max_in_flight_packets: usize, timeout: Duration) -> Self {
Self {
partial: HashMap::new(),
recent: HashSet::new(),
recent_order: VecDeque::new(),
max_in_flight_packets: max_in_flight_packets.max(1),
timeout_ms: u64::try_from(timeout.as_millis()).unwrap_or(u64::MAX),
}
}
pub fn push(
&mut self,
frame: CarrierFrame,
now_ms: u64,
) -> Result<Option<Vec<u8>>, CarrierFrameError> {
self.expire(now_ms);
let packet_id = frame.header.packet_id;
if self.recent.contains(&packet_id) {
return Err(CarrierFrameError::DuplicatePacket(packet_id));
}
if !self.partial.contains_key(&packet_id)
&& self.partial.len() >= self.max_in_flight_packets
{
return Err(CarrierFrameError::Backpressure);
}
let segment_count = usize::from(frame.header.segment_count);
let packet = self
.partial
.entry(packet_id)
.or_insert_with(|| PartialPacket {
created_at_ms: now_ms,
segments: vec![None; segment_count],
total_bytes: 0,
});
if packet.segments.len() != segment_count {
self.partial.remove(&packet_id);
return Err(CarrierFrameError::SegmentCountChanged(packet_id));
}
let index = usize::from(frame.header.segment_index);
if packet.segments[index].is_some() {
return Err(CarrierFrameError::DuplicateSegment {
packet_id,
segment_index: frame.header.segment_index,
});
}
packet.total_bytes = packet
.total_bytes
.checked_add(frame.payload.len())
.ok_or(CarrierFrameError::PacketSize(usize::MAX))?;
if packet.total_bytes > MAX_INNER_PACKET_BYTES {
let total_bytes = packet.total_bytes;
self.partial.remove(&packet_id);
return Err(CarrierFrameError::PacketSize(total_bytes));
}
packet.segments[index] = Some(frame.payload);
if packet.segments.iter().any(Option::is_none) {
return Ok(None);
}
let packet = self
.partial
.remove(&packet_id)
.expect("packet still present");
let mut output = Vec::with_capacity(packet.total_bytes);
for segment in packet.segments {
output.extend(segment.expect("completion checked"));
}
self.record_completed(packet_id);
Ok(Some(output))
}
pub fn expire(&mut self, now_ms: u64) {
self.partial
.retain(|_, packet| now_ms.saturating_sub(packet.created_at_ms) < self.timeout_ms);
}
fn record_completed(&mut self, packet_id: u64) {
self.recent.insert(packet_id);
self.recent_order.push_back(packet_id);
while self.recent_order.len() > RECENT_PACKET_IDS {
if let Some(retired) = self.recent_order.pop_front() {
self.recent.remove(&retired);
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CarrierFrameError {
TruncatedHeader { actual: usize },
WrongMagic,
UnsupportedVersion(u8),
WrongTransport { expected: u64, actual: u64 },
WrongSession,
StaleGeneration { expected: u64, actual: u64 },
InvalidSegmentCount(u16),
InvalidSegmentIndex { index: u16, count: u16 },
PayloadLength { declared: usize, actual: usize },
PacketSize(usize),
CarrierMtu(usize),
TooManySegments,
SegmentCountChanged(u64),
DuplicatePacket(u64),
DuplicateSegment { packet_id: u64, segment_index: u16 },
Backpressure,
InvalidControl,
InvalidControlAuthentication,
ReservedPacketId,
}
impl std::fmt::Display for CarrierFrameError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(formatter, "{self:?}")
}
}
impl std::error::Error for CarrierFrameError {}
#[cfg(test)]
mod tests {
use super::*;
fn expected() -> CarrierFrameExpectation {
CarrierFrameExpectation {
transport_id: 0x57_52_54_43,
carrier_session_id: [7; 16],
transport_generation: 41,
}
}
#[test]
fn frame_round_trip_validates_every_identity_domain() {
let frames = segment_packet(b"iroh packet", expected(), 9, 256).unwrap();
let encoded = frames[0].encode().unwrap();
let decoded = CarrierFrame::decode(&encoded, expected()).unwrap();
assert_eq!(decoded, frames[0]);
let mut stale = expected();
stale.transport_generation += 1;
assert!(matches!(
CarrierFrame::decode(&encoded, stale),
Err(CarrierFrameError::StaleGeneration { .. })
));
}
#[test]
fn segmented_packet_reassembles_out_of_order_once() {
let packet = vec![0x5a; MAX_INNER_PACKET_BYTES];
let mut frames = segment_packet(&packet, expected(), 77, 128).unwrap();
frames.reverse();
let now_ms = 10;
let mut reassembler = CarrierReassembler::default();
let mut completed = None;
for frame in frames {
completed = reassembler.push(frame, now_ms).unwrap().or(completed);
}
assert_eq!(completed.as_deref(), Some(packet.as_slice()));
let replay = segment_packet(&packet, expected(), 77, 128)
.unwrap()
.remove(0);
assert_eq!(
reassembler.push(replay, now_ms),
Err(CarrierFrameError::DuplicatePacket(77))
);
}
#[test]
fn malformed_segment_metadata_is_rejected_before_payload() {
let mut frame = segment_packet(b"packet", expected(), 1, 128)
.unwrap()
.remove(0);
frame.header.segment_count = 0;
assert_eq!(
frame.encode(),
Err(CarrierFrameError::InvalidSegmentCount(0))
);
}
#[test]
fn reassembly_is_bounded_and_expires() {
let now_ms = 10;
let mut reassembler = CarrierReassembler::new(1, Duration::from_millis(5));
let first = segment_packet(&vec![1; 100], expected(), 1, 80)
.unwrap()
.remove(0);
let second = segment_packet(&vec![2; 100], expected(), 2, 80)
.unwrap()
.remove(0);
assert_eq!(reassembler.push(first, now_ms), Ok(None));
assert_eq!(
reassembler.push(second.clone(), now_ms),
Err(CarrierFrameError::Backpressure)
);
assert_eq!(reassembler.push(second, now_ms + 6), Ok(None));
}
#[test]
fn packet_and_carrier_mtu_limits_fail_closed() {
assert_eq!(
segment_packet(&vec![0; MAX_INNER_PACKET_BYTES + 1], expected(), 1, 256),
Err(CarrierFrameError::PacketSize(MAX_INNER_PACKET_BYTES + 1))
);
assert_eq!(
segment_packet(b"x", expected(), 1, CARRIER_HEADER_LEN),
Err(CarrierFrameError::CarrierMtu(CARRIER_HEADER_LEN))
);
}
#[test]
fn terminal_control_is_generation_fenced_and_authenticated() {
let key = [0x5a; 32];
let control =
CarrierFrame::terminal_control(expected(), &key, CarrierControl::SessionTokenRevoked);
let encoded = control.encode().unwrap();
let decoded = CarrierFrame::decode(&encoded, expected()).unwrap();
assert_eq!(
decoded.terminal_control_kind(expected(), &key),
Ok(Some(CarrierControl::SessionTokenRevoked))
);
assert_eq!(
decoded.terminal_control_kind(expected(), &[0x33; 32]),
Err(CarrierFrameError::InvalidControlAuthentication)
);
let mut stale = expected();
stale.transport_generation += 1;
assert!(matches!(
CarrierFrame::decode(&encoded, stale),
Err(CarrierFrameError::StaleGeneration { .. })
));
}
}