use serde::Serialize;
use serde_json::{Map, Value};
use std::fmt;
pub const CHANNEL_ENVELOPE_MAGIC: u8 = 0x7f;
pub const CHANNEL_ENVELOPE_VERSION: u8 = 0x01;
pub const CHANNEL_ENVELOPE_HEADER_BYTES: usize = 1 + 1 + 2 + 2;
pub const EXPLICIT_PROTOCOL_BYTE: u8 = 0x02;
pub const NATIVE_MAIN_PROTOCOL_BYTE: u8 = 0x00;
pub const NATIVE_MAIN_LABEL: &str = "main";
pub const NATIVE_MAIN_MAX_LABEL_BYTES: usize = 64;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "kebab-case")]
pub enum IncomingProtocolHint {
Control,
Explicit,
NativeMain,
Unknown,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct IncomingStreamChannelMetadata {
pub channel_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<Map<String, Value>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IncomingStreamInspection {
pub protocol_hint: IncomingProtocolHint,
pub channel: Option<IncomingStreamChannelMetadata>,
pub consumed_prefix_len: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum IncomingStreamInspectionDecision {
NeedMore(usize),
Complete(IncomingStreamInspection),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ChannelEnvelopePrefix {
pub channel: IncomingStreamChannelMetadata,
pub consumed_len: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ChannelEnvelopeDecodeDecision {
NeedMore(usize),
NotMatched,
Decoded(ChannelEnvelopePrefix),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ChannelEnvelopeEncodeError {
EmptyChannelId,
ChannelIdTooLarge(usize),
MetadataTooLarge(usize),
MetadataSerialize(String),
}
impl fmt::Display for ChannelEnvelopeEncodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::EmptyChannelId => write!(f, "channel envelope requires a non-empty channel id"),
Self::ChannelIdTooLarge(len) => {
write!(f, "channel envelope channel id is too large: {len} bytes")
}
Self::MetadataTooLarge(len) => {
write!(f, "channel envelope metadata is too large: {len} bytes")
}
Self::MetadataSerialize(error) => {
write!(f, "channel envelope metadata failed to serialize: {error}")
}
}
}
}
impl std::error::Error for ChannelEnvelopeEncodeError {}
#[derive(Debug, Clone, PartialEq, Eq)]
enum PrefixParseResult<T> {
NeedMore(usize),
NotMatched,
Parsed(T),
}
pub fn encode_channel_envelope(
channel_id: &str,
metadata: Option<&Map<String, Value>>,
) -> Result<Vec<u8>, ChannelEnvelopeEncodeError> {
let channel_id = channel_id.trim();
if channel_id.is_empty() {
return Err(ChannelEnvelopeEncodeError::EmptyChannelId);
}
let channel_bytes = channel_id.as_bytes();
if channel_bytes.len() > u16::MAX as usize {
return Err(ChannelEnvelopeEncodeError::ChannelIdTooLarge(
channel_bytes.len(),
));
}
let metadata_bytes = metadata
.map(|value| {
serde_json::to_vec(value)
.map_err(|error| ChannelEnvelopeEncodeError::MetadataSerialize(error.to_string()))
})
.transpose()?
.unwrap_or_default();
if metadata_bytes.len() > u16::MAX as usize {
return Err(ChannelEnvelopeEncodeError::MetadataTooLarge(
metadata_bytes.len(),
));
}
let mut bytes = Vec::with_capacity(
CHANNEL_ENVELOPE_HEADER_BYTES + channel_bytes.len() + metadata_bytes.len(),
);
bytes.push(CHANNEL_ENVELOPE_MAGIC);
bytes.push(CHANNEL_ENVELOPE_VERSION);
bytes.extend_from_slice(&(channel_bytes.len() as u16).to_be_bytes());
bytes.extend_from_slice(&(metadata_bytes.len() as u16).to_be_bytes());
bytes.extend_from_slice(channel_bytes);
bytes.extend_from_slice(&metadata_bytes);
Ok(bytes)
}
pub fn decode_channel_envelope_prefix(bytes: &[u8]) -> ChannelEnvelopeDecodeDecision {
match parse_channel_envelope(bytes) {
PrefixParseResult::NeedMore(required) => ChannelEnvelopeDecodeDecision::NeedMore(required),
PrefixParseResult::NotMatched => ChannelEnvelopeDecodeDecision::NotMatched,
PrefixParseResult::Parsed(prefix) => ChannelEnvelopeDecodeDecision::Decoded(prefix),
}
}
pub fn inspect_incoming_stream_prefix(
bytes: &[u8],
reached_eof: bool,
) -> IncomingStreamInspectionDecision {
let channel_prefix = match parse_channel_envelope(bytes) {
PrefixParseResult::Parsed(channel) => Some(channel),
PrefixParseResult::NeedMore(required) if !reached_eof => {
return IncomingStreamInspectionDecision::NeedMore(required);
}
PrefixParseResult::NeedMore(_) | PrefixParseResult::NotMatched => None,
};
let channel = channel_prefix.as_ref().map(|value| value.channel.clone());
let consumed_prefix_len = channel_prefix
.as_ref()
.map(|value| value.consumed_len)
.unwrap_or(0);
let remaining = &bytes[consumed_prefix_len..];
if remaining.is_empty() {
let protocol_hint = if reached_eof {
IncomingProtocolHint::Unknown
} else {
return IncomingStreamInspectionDecision::NeedMore(consumed_prefix_len + 1);
};
return IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint,
channel,
consumed_prefix_len,
});
}
let first_byte = remaining[0];
if first_byte == EXPLICIT_PROTOCOL_BYTE {
return IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint: IncomingProtocolHint::Explicit,
channel,
consumed_prefix_len,
});
}
let protocol_hint = match parse_native_main_prefix(remaining) {
PrefixParseResult::Parsed(native_main_prefix_len) => {
return IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint: IncomingProtocolHint::NativeMain,
channel,
consumed_prefix_len: consumed_prefix_len + native_main_prefix_len,
});
}
PrefixParseResult::NeedMore(required) if !reached_eof => {
return IncomingStreamInspectionDecision::NeedMore(consumed_prefix_len + required);
}
PrefixParseResult::NeedMore(_) | PrefixParseResult::NotMatched => {
IncomingProtocolHint::Control
}
};
IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint,
channel,
consumed_prefix_len,
})
}
fn parse_channel_envelope(bytes: &[u8]) -> PrefixParseResult<ChannelEnvelopePrefix> {
if bytes.is_empty() {
return PrefixParseResult::NeedMore(1);
}
if bytes[0] != CHANNEL_ENVELOPE_MAGIC {
return PrefixParseResult::NotMatched;
}
if bytes.len() < 2 {
return PrefixParseResult::NeedMore(2);
}
if bytes[1] != CHANNEL_ENVELOPE_VERSION {
return PrefixParseResult::NotMatched;
}
if bytes.len() < CHANNEL_ENVELOPE_HEADER_BYTES {
return PrefixParseResult::NeedMore(CHANNEL_ENVELOPE_HEADER_BYTES);
}
let channel_len = u16::from_be_bytes([bytes[2], bytes[3]]) as usize;
let metadata_len = u16::from_be_bytes([bytes[4], bytes[5]]) as usize;
let total_envelope_bytes = CHANNEL_ENVELOPE_HEADER_BYTES + channel_len + metadata_len;
if channel_len == 0 {
return PrefixParseResult::NotMatched;
}
if bytes.len() < total_envelope_bytes {
return PrefixParseResult::NeedMore(total_envelope_bytes);
}
let channel_start = CHANNEL_ENVELOPE_HEADER_BYTES;
let channel_end = channel_start + channel_len;
let metadata_end = channel_end + metadata_len;
let channel_id = String::from_utf8_lossy(&bytes[channel_start..channel_end])
.trim()
.to_string();
if channel_id.is_empty() {
return PrefixParseResult::NotMatched;
}
let metadata = if metadata_len == 0 {
None
} else {
match serde_json::from_slice::<Value>(&bytes[channel_end..metadata_end]) {
Ok(Value::Object(map)) => Some(map),
_ => None,
}
};
PrefixParseResult::Parsed(ChannelEnvelopePrefix {
channel: IncomingStreamChannelMetadata {
channel_id,
metadata,
},
consumed_len: total_envelope_bytes,
})
}
fn parse_native_main_prefix(bytes: &[u8]) -> PrefixParseResult<usize> {
if bytes.is_empty() {
return PrefixParseResult::NeedMore(1);
}
if bytes[0] != NATIVE_MAIN_PROTOCOL_BYTE {
return PrefixParseResult::NotMatched;
}
if bytes.len() < 5 {
return PrefixParseResult::NeedMore(5);
}
let label_len = u32::from_be_bytes([bytes[1], bytes[2], bytes[3], bytes[4]]) as usize;
if label_len == 0 || label_len > NATIVE_MAIN_MAX_LABEL_BYTES {
return PrefixParseResult::NotMatched;
}
let total_prefix_bytes = 5 + label_len;
if bytes.len() < total_prefix_bytes {
return PrefixParseResult::NeedMore(total_prefix_bytes);
}
let label = String::from_utf8_lossy(&bytes[5..total_prefix_bytes]);
if label != NATIVE_MAIN_LABEL {
return PrefixParseResult::NotMatched;
}
PrefixParseResult::Parsed(total_prefix_bytes)
}
#[cfg(test)]
mod tests {
use super::*;
fn encode_native_main_prefix() -> Vec<u8> {
let label = NATIVE_MAIN_LABEL.as_bytes();
let mut bytes = Vec::with_capacity(1 + 4 + label.len());
bytes.push(NATIVE_MAIN_PROTOCOL_BYTE);
bytes.extend_from_slice(&(label.len() as u32).to_be_bytes());
bytes.extend_from_slice(label);
bytes
}
#[test]
fn inspects_explicit_streams_without_consuming_protocol_byte() {
let bytes = vec![EXPLICIT_PROTOCOL_BYTE, 0xaa, 0xbb];
let inspected = inspect_incoming_stream_prefix(&bytes, false);
assert_eq!(
inspected,
IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint: IncomingProtocolHint::Explicit,
channel: None,
consumed_prefix_len: 0,
})
);
}
#[test]
fn inspects_channel_envelope_and_preserves_explicit_payload() {
let metadata = serde_json::json!({ "scope": "ticket-only" })
.as_object()
.cloned()
.expect("object metadata");
let mut bytes =
encode_channel_envelope("share-control", Some(&metadata)).expect("channel envelope");
bytes.extend_from_slice(&[EXPLICIT_PROTOCOL_BYTE, 0x01, 0x02]);
let inspected = inspect_incoming_stream_prefix(&bytes, false);
assert_eq!(
inspected,
IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint: IncomingProtocolHint::Explicit,
channel: Some(IncomingStreamChannelMetadata {
channel_id: "share-control".to_string(),
metadata: Some(
serde_json::json!({ "scope": "ticket-only" })
.as_object()
.cloned()
.expect("object metadata"),
),
}),
consumed_prefix_len: CHANNEL_ENVELOPE_HEADER_BYTES
+ "share-control".len()
+ serde_json::to_string(&serde_json::json!({ "scope": "ticket-only" }))
.expect("json string")
.len(),
})
);
}
#[test]
fn inspects_native_main_and_consumes_label_prefix_only() {
let mut bytes = encode_native_main_prefix();
bytes.extend_from_slice(&[0x00, 0x00, 0x00, 0x10]);
let inspected = inspect_incoming_stream_prefix(&bytes, false);
assert_eq!(
inspected,
IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint: IncomingProtocolHint::NativeMain,
channel: None,
consumed_prefix_len: 1 + 4 + NATIVE_MAIN_LABEL.len(),
})
);
}
#[test]
fn incomplete_native_main_prefix_falls_back_to_control_at_eof() {
let bytes = vec![NATIVE_MAIN_PROTOCOL_BYTE, 0x00];
let inspected = inspect_incoming_stream_prefix(&bytes, true);
assert_eq!(
inspected,
IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint: IncomingProtocolHint::Control,
channel: None,
consumed_prefix_len: 0,
})
);
}
#[test]
fn parsed_channel_with_no_remaining_payload_is_unknown_at_eof() {
let bytes = encode_channel_envelope("share", None).expect("channel envelope");
let inspected = inspect_incoming_stream_prefix(&bytes, true);
assert_eq!(
inspected,
IncomingStreamInspectionDecision::Complete(IncomingStreamInspection {
protocol_hint: IncomingProtocolHint::Unknown,
channel: Some(IncomingStreamChannelMetadata {
channel_id: "share".to_string(),
metadata: None,
}),
consumed_prefix_len: CHANNEL_ENVELOPE_HEADER_BYTES + "share".len(),
})
);
}
#[test]
fn encodes_latency_probe_channel_envelope_golden_vector() {
let bytes =
encode_channel_envelope("system/device-latency-probe", None).expect("channel envelope");
assert_eq!(
bytes,
hex_to_bytes("7f01001b000073797374656d2f6465766963652d6c6174656e63792d70726f6265"),
);
}
#[test]
fn decodes_channel_envelope_prefix_directly() {
let metadata = serde_json::json!({ "scope": "ticket-only" })
.as_object()
.cloned()
.expect("object metadata");
let bytes =
encode_channel_envelope("share-control", Some(&metadata)).expect("channel envelope");
assert_eq!(
decode_channel_envelope_prefix(&bytes),
ChannelEnvelopeDecodeDecision::Decoded(ChannelEnvelopePrefix {
channel: IncomingStreamChannelMetadata {
channel_id: "share-control".to_string(),
metadata: Some(metadata),
},
consumed_len: bytes.len(),
})
);
}
fn hex_to_bytes(value: &str) -> Vec<u8> {
value
.as_bytes()
.chunks(2)
.map(|chunk| {
let text = std::str::from_utf8(chunk).expect("hex utf8");
u8::from_str_radix(text, 16).expect("hex byte")
})
.collect()
}
}