use std::fmt;
use std::io::Cursor;
use binrw::{BinRead, BinWrite};
use thiserror::Error;
use super::catalog::MessageId;
use super::catalog::MessageRoute;
pub const HEADER_SIZE: usize = 12;
pub const MAX_FRAME_SIZE: usize = 8192;
#[derive(BinRead, BinWrite, Clone, Copy, Debug, Eq, PartialEq)]
#[brw(little)]
struct WireHeader {
wire_len: u32,
protocol_version: u32,
message_id: u32,
}
#[derive(Clone, Eq, PartialEq)]
pub struct Frame {
pub protocol_version: u32,
pub message_id: u32,
pub payload: Vec<u8>,
}
impl fmt::Debug for Frame {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Frame")
.field("protocol_version", &self.protocol_version)
.field("message_id", &format_args!("0x{:04x}", self.message_id))
.field("payload_len", &self.payload.len())
.finish()
}
}
impl Frame {
pub fn new(protocol_version: u32, message_id: u32, payload: Vec<u8>) -> Self {
Self {
protocol_version,
message_id,
payload,
}
}
pub fn encode(&self) -> Result<Vec<u8>, CodecError> {
let wire_len = self
.payload
.len()
.checked_add(4)
.ok_or(CodecError::FrameTooLarge(usize::MAX))?;
let total = wire_len + 8;
if total > MAX_FRAME_SIZE {
return Err(CodecError::FrameTooLarge(total));
}
let header = WireHeader {
wire_len: u32::try_from(wire_len).map_err(|_| CodecError::FrameTooLarge(wire_len))?,
protocol_version: self.protocol_version,
message_id: self.message_id,
};
let mut output = Cursor::new(Vec::with_capacity(total));
header
.write(&mut output)
.map_err(|error| CodecError::wire("encode", self.message_id, &error))?;
output.get_mut().extend_from_slice(&self.payload);
Ok(output.into_inner())
}
pub fn message_type(&self) -> MessageId {
MessageId::from(self.message_id)
}
}
#[derive(Debug, Error, Clone, Eq, PartialEq)]
pub enum CodecError {
#[error("SCCP frame length {0} is invalid")]
InvalidLength(u32),
#[error("SCCP frame size {0} exceeds the configured maximum")]
FrameTooLarge(usize),
#[error("message 0x{message_id:04x} is truncated: need {needed} bytes, got {actual}")]
Truncated {
message_id: u32,
needed: usize,
actual: usize,
},
#[error("invalid SCCP device ID: {0}")]
InvalidDeviceId(String),
#[error("invalid SCCP device definition: {0}")]
InvalidDefinition(String),
#[error("invalid UTF-8-compatible SCCP text field")]
InvalidText,
#[error("unsupported SCCP protocol version {0}")]
UnsupportedProtocol(u32),
#[error("message 0x{message_id:04x} has route {actual:?}, expected {expected}")]
UnexpectedRoute {
message_id: u32,
actual: MessageRoute,
expected: &'static str,
},
#[error("message 0x{message_id:04x} contains invalid {field}: {value}")]
InvalidValue {
message_id: u32,
field: &'static str,
value: u64,
},
#[error("message 0x{message_id:04x} {field} count {count} exceeds maximum {maximum}")]
CountTooLarge {
message_id: u32,
field: &'static str,
count: usize,
maximum: usize,
},
#[error("message 0x{message_id:04x} contains non-zero {field} padding")]
NonZeroPadding {
message_id: u32,
field: &'static str,
},
#[error("message 0x{message_id:04x} has {count} unexpected trailing bytes")]
TrailingBytes { message_id: u32, count: usize },
#[error(
"message 0x{message_id:04x} payload length {actual} is not aligned to a four-byte boundary"
)]
InvalidAlignment { message_id: u32, actual: usize },
#[error(
"message 0x{message_id:04x} field {field} is too long: {actual} bytes, maximum {maximum}"
)]
TextTooLong {
message_id: u32,
field: &'static str,
actual: usize,
maximum: usize,
},
#[error("secret field {field} is too long: {actual} bytes, maximum {maximum}")]
SecretTooLong {
field: &'static str,
actual: usize,
maximum: usize,
},
#[error("could not {operation} SCCP message 0x{message_id:04x} at byte {offset}: {detail}")]
Wire {
operation: &'static str,
message_id: u32,
offset: u64,
detail: String,
},
}
impl CodecError {
pub(crate) fn wire(operation: &'static str, message_id: u32, error: &binrw::Error) -> Self {
let offset = match error {
binrw::Error::BadMagic { pos, .. }
| binrw::Error::AssertFail { pos, .. }
| binrw::Error::Custom { pos, .. }
| binrw::Error::NoVariantMatch { pos }
| binrw::Error::EnumErrors { pos, .. } => *pos,
binrw::Error::Io(_) | binrw::Error::Backtrace(_) => 0,
_ => 0,
};
Self::Wire {
operation,
message_id,
offset,
detail: error.to_string(),
}
}
}
#[derive(Debug, Default)]
pub struct FrameDecoder {
buffer: Vec<u8>,
}
impl FrameDecoder {
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, bytes: &[u8]) -> Result<Vec<Frame>, CodecError> {
self.buffer.extend_from_slice(bytes);
let mut frames = Vec::new();
let mut consumed = 0_usize;
loop {
let retained = &self.buffer[consumed..];
if retained.len() < HEADER_SIZE {
break;
}
let header = WireHeader::read(&mut Cursor::new(&retained[..HEADER_SIZE]))
.map_err(|error| CodecError::wire("decode header for", 0, &error))?;
let length = header.wire_len;
if length < 4 {
return Err(CodecError::InvalidLength(length));
}
let total =
usize::try_from(length).map_err(|_| CodecError::FrameTooLarge(usize::MAX))? + 8;
if total > MAX_FRAME_SIZE {
return Err(CodecError::FrameTooLarge(total));
}
if retained.len() < total {
break;
}
let payload = retained[HEADER_SIZE..total].to_vec();
consumed += total;
frames.push(Frame {
protocol_version: header.protocol_version,
message_id: header.message_id,
payload,
});
}
if consumed != 0 {
self.buffer.drain(..consumed);
}
debug_assert!(self.buffer.len() < MAX_FRAME_SIZE);
Ok(frames)
}
pub fn buffered_len(&self) -> usize {
self.buffer.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn one_large_chunk_of_small_frames_is_drained_incrementally() {
let encoded = Frame::new(22, 0x0100, Vec::new()).encode().unwrap();
let frame_count = MAX_FRAME_SIZE * 2 / encoded.len() + 1;
let stream = encoded.repeat(frame_count);
assert!(stream.len() > MAX_FRAME_SIZE * 2);
let frames = FrameDecoder::new().push(&stream).unwrap();
assert_eq!(frames.len(), frame_count);
assert!(frames.iter().all(|frame| {
frame.protocol_version == 22 && frame.message_id == 0x0100 && frame.payload.is_empty()
}));
}
#[test]
fn frame_length_matches_skinny_header_definition() {
let frame = Frame::new(22, 0x0006, vec![1, 2, 3, 4]);
let encoded = frame.encode().unwrap();
assert_eq!(&encoded[..4], &8_u32.to_le_bytes());
assert_eq!(encoded.len(), 16);
}
#[test]
fn decoder_handles_fragmented_and_coalesced_frames() {
let first = Frame::new(0, 0, Vec::new()).encode().unwrap();
let second = Frame::new(22, 6, vec![1; 8]).encode().unwrap();
let mut decoder = FrameDecoder::new();
assert!(decoder.push(&first[..5]).unwrap().is_empty());
let mut rest = first[5..].to_vec();
rest.extend_from_slice(&second);
rest.extend_from_slice(&first);
rest.extend_from_slice(&second);
let frames = decoder.push(&rest).unwrap();
assert_eq!(
frames
.iter()
.map(|frame| frame.message_id)
.collect::<Vec<_>>(),
[0, 6, 0, 6]
);
assert_eq!(frames[1], frames[3], "duplicate frames changed in transit");
assert_eq!(frames[0], frames[2], "reordered frames changed in transit");
}
#[test]
fn decoder_retains_only_the_incomplete_tail_after_many_frames() {
let complete = Frame::new(22, 0x0100, vec![1, 2, 3, 4]).encode().unwrap();
let tail = Frame::new(22, 0x0101, vec![5; 32]).encode().unwrap();
let split = tail.len() - 7;
let mut chunk = complete.repeat(1_000);
chunk.extend_from_slice(&tail[..split]);
let mut decoder = FrameDecoder::new();
assert_eq!(decoder.push(&chunk).unwrap().len(), 1_000);
assert_eq!(decoder.buffered_len(), split);
let frames = decoder.push(&tail[split..]).unwrap();
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].message_id, 0x0101);
assert_eq!(decoder.buffered_len(), 0);
}
#[test]
fn decoder_accepts_every_possible_single_fragment_boundary() {
let bytes = Frame::new(22, 0x22, (0_u8..64).collect()).encode().unwrap();
for split in 0..bytes.len() {
let mut decoder = FrameDecoder::new();
assert!(decoder.push(&bytes[..split]).unwrap().is_empty());
let frames = decoder.push(&bytes[split..]).unwrap();
assert_eq!(frames.len(), 1, "split at byte {split}");
assert_eq!(frames[0].payload, (0_u8..64).collect::<Vec<_>>());
assert_eq!(decoder.buffered_len(), 0);
}
}
#[test]
fn invalid_short_length_is_rejected() {
let mut decoder = FrameDecoder::new();
let mut bytes = vec![0; 12];
bytes[..4].copy_from_slice(&3_u32.to_le_bytes());
assert_eq!(decoder.push(&bytes), Err(CodecError::InvalidLength(3)));
}
}