1use crate::constants::{MAGIC, MAX_PAYLOAD_SIZE, VERSION};
8use crate::error::ProtocolError;
9use crate::frame::{Frame, FrameHeader};
10use crate::types::{FrameFlags, MessageType, ProtocolErrorKind, RequestId};
11
12pub fn encode(
21 msg_type: MessageType,
22 flags: FrameFlags,
23 request_id: RequestId,
24 payload: &[u8],
25) -> Result<Vec<u8>, ProtocolError> {
26 if payload.len() > MAX_PAYLOAD_SIZE {
27 return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
28 }
29
30 let mut buf = Vec::with_capacity(FrameHeader::SIZE + payload.len());
31
32 buf.extend_from_slice(&MAGIC.to_le_bytes());
34 buf.extend_from_slice(&VERSION.to_le_bytes());
35 buf.push(msg_type as u8);
36 buf.push(flags.bits());
37 buf.extend_from_slice(&request_id.0.to_le_bytes());
38 buf.extend_from_slice(&u32::try_from(payload.len()).unwrap().to_le_bytes());
40
41 buf.extend_from_slice(payload);
43
44 debug_assert_eq!(buf.len(), FrameHeader::SIZE + payload.len());
45
46 Ok(buf)
47}
48
49pub fn decode(buffer: &[u8]) -> Result<Frame<'_>, ProtocolError> {
68 if buffer.len() < FrameHeader::SIZE {
69 return Err(ProtocolError::new(ProtocolErrorKind::MalformedFrame));
70 }
71
72 let magic = u32::from_le_bytes(buffer[0..4].try_into().unwrap());
74 if magic != MAGIC {
75 return Err(ProtocolError::new(ProtocolErrorKind::InvalidMagic));
76 }
77
78 let version = u16::from_le_bytes(buffer[4..6].try_into().unwrap());
79 if version != VERSION {
80 return Err(ProtocolError::new(ProtocolErrorKind::UnsupportedVersion));
81 }
82
83 let msg_type = MessageType::try_from(buffer[6])
84 .map_err(|()| ProtocolError::new(ProtocolErrorKind::UnknownMessageType))?;
85
86 let flags = FrameFlags::from_bits_truncate(buffer[7]);
87
88 let request_id = RequestId(u64::from_le_bytes(buffer[8..16].try_into().unwrap()));
89
90 let payload_length_u32 = u32::from_le_bytes(buffer[16..20].try_into().unwrap());
91 let payload_length = payload_length_u32 as usize;
92 if payload_length > MAX_PAYLOAD_SIZE {
93 return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
94 }
95
96 let total_len = FrameHeader::SIZE + payload_length;
97 if buffer.len() < total_len {
98 return Err(ProtocolError::new(ProtocolErrorKind::MalformedFrame));
99 }
100
101 let payload = &buffer[FrameHeader::SIZE..total_len];
102
103 Ok(Frame {
104 header: FrameHeader {
105 magic,
106 version,
107 msg_type: msg_type as u8,
108 flags: flags.bits(),
109 request_id: request_id.0,
110 payload_length: payload_length_u32,
111 },
112 payload,
113 })
114}