use bytes::{Buf, BufMut, Bytes, BytesMut};
use derive_builder::Builder;
use std::collections::HashMap;
use thiserror::Error;
use super::responses::ResponseId;
const CURRENT_SCHEMA_VERSION: u8 = 1;
const MAX_HEADER_VALUE_LEN: usize = 1024;
const MAX_HEADERS_LEN: usize = 16384;
#[derive(Debug, Clone)]
pub(crate) struct ActiveMessage {
pub metadata: MessageMetadata,
pub payload: Bytes,
}
impl ActiveMessage {
pub(crate) fn encode(
self,
) -> Result<(Bytes, Bytes, crate::transports::MessageType), EncodeError> {
encode_active_message(self)
}
}
#[derive(Debug, Clone, Builder)]
#[builder(setter(into))]
pub(crate) struct MessageMetadata {
#[builder(default = "CURRENT_SCHEMA_VERSION")]
pub schema_version: u8,
pub response_type: ResponseType,
pub response_id: ResponseId,
pub handler_name: String,
#[builder(default)]
pub headers: Option<HashMap<String, String>>,
}
impl MessageMetadata {
pub(crate) fn new_fire(
response_id: ResponseId,
handler_name: String,
headers: Option<HashMap<String, String>>,
) -> Self {
Self {
schema_version: CURRENT_SCHEMA_VERSION,
response_type: ResponseType::FireAndForget,
response_id,
handler_name,
headers,
}
}
pub(crate) fn new_sync(
response_id: ResponseId,
handler_name: String,
headers: Option<HashMap<String, String>>,
) -> Self {
Self {
schema_version: CURRENT_SCHEMA_VERSION,
response_type: ResponseType::AckNack,
response_id,
handler_name,
headers,
}
}
pub(crate) fn new_unary(
response_id: ResponseId,
handler_name: String,
headers: Option<HashMap<String, String>>,
) -> Self {
Self {
schema_version: CURRENT_SCHEMA_VERSION,
response_type: ResponseType::Unary,
response_id,
handler_name,
headers,
}
}
}
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ResponseType {
FireAndForget = 0,
AckNack = 1,
Unary = 2,
}
impl TryFrom<u8> for ResponseType {
type Error = DecodeError;
fn try_from(value: u8) -> Result<Self, Self::Error> {
match value {
0 => Ok(ResponseType::FireAndForget),
1 => Ok(ResponseType::AckNack),
2 => Ok(ResponseType::Unary),
_ => Err(DecodeError::InvalidResponseType(value)),
}
}
}
impl ResponseType {
pub(crate) fn to_message_type(self) -> crate::transports::MessageType {
crate::transports::MessageType::Message
}
}
#[derive(Debug, Error)]
pub(crate) enum DecodeError {
#[error("Header too short: expected at least 20 bytes")]
HeaderTooShort,
#[error("Invalid handler name length")]
InvalidHandlerNameLength,
#[error("Invalid UTF-8 in handler name")]
InvalidUtf8,
#[error("Invalid response type: {0}")]
InvalidResponseType(u8),
#[error("Invalid headers length")]
InvalidHeadersLength,
#[error("Unsupported schema version: got {0}, expected {1}")]
UnsupportedSchemaVersion(u8, u8),
#[error("Failed to deserialize headers: {0}")]
HeaderDeserializationError(#[from] rmp_serde::decode::Error),
}
#[derive(Debug, Error)]
pub(crate) enum EncodeError {
#[error("Handler name too long: {0} bytes exceeds maximum of 65535")]
HandlerNameTooLong(usize),
#[error("Header value too large: key '{0}' has value of {1} bytes, max is 1024")]
HeaderValueTooLarge(String, usize),
#[error("Total headers too large: {0} bytes exceeds maximum of 16384")]
TotalHeadersTooLarge(usize),
#[error("Failed to serialize headers: {0}")]
HeaderSerializationError(#[from] rmp_serde::encode::Error),
}
pub(crate) fn encode_active_message(
message: ActiveMessage,
) -> Result<(Bytes, Bytes, crate::transports::MessageType), EncodeError> {
let handler_name_len = message.metadata.handler_name.len();
if handler_name_len > u16::MAX as usize {
return Err(EncodeError::HandlerNameTooLong(handler_name_len));
}
let headers_bytes = if let Some(ref headers) = message.metadata.headers {
for (key, value) in headers.iter() {
if value.len() > MAX_HEADER_VALUE_LEN {
return Err(EncodeError::HeaderValueTooLarge(key.clone(), value.len()));
}
}
let msgpack_bytes = rmp_serde::to_vec(headers)?;
if msgpack_bytes.len() > MAX_HEADERS_LEN {
return Err(EncodeError::TotalHeadersTooLarge(msgpack_bytes.len()));
}
Some(msgpack_bytes)
} else {
None
};
let headers_len = headers_bytes.as_ref().map(|b| b.len()).unwrap_or(0);
let header_size = 20 + handler_name_len + 2 + headers_len; let mut header = BytesMut::with_capacity(header_size);
header.put_u8(message.metadata.schema_version);
header.put_u8(message.metadata.response_type as u8);
header.put_u128_le(message.metadata.response_id.as_u128());
header.put_u16_le(handler_name_len as u16);
header.put_slice(message.metadata.handler_name.as_bytes());
header.put_u16_le(headers_len as u16);
if let Some(bytes) = headers_bytes {
header.put_slice(&bytes);
}
let message_type = message.metadata.response_type.to_message_type();
Ok((header.freeze(), message.payload, message_type))
}
pub(crate) fn decode_response_id_from_request_header(header: &Bytes) -> Option<ResponseId> {
if header.len() < 18 {
return None;
}
if header[0] != CURRENT_SCHEMA_VERSION {
return None;
}
if ResponseType::try_from(header[1]).is_err() {
return None;
}
let mut id_bytes = [0u8; 16];
id_bytes.copy_from_slice(&header[2..18]);
Some(ResponseId::from_u128(u128::from_le_bytes(id_bytes)))
}
pub(crate) fn decode_active_message(
header: Bytes,
payload: Bytes,
) -> Result<ActiveMessage, DecodeError> {
let mut header = header;
if header.len() < 22 {
return Err(DecodeError::HeaderTooShort);
}
let schema_version = header.get_u8();
if schema_version != CURRENT_SCHEMA_VERSION {
return Err(DecodeError::UnsupportedSchemaVersion(
schema_version,
CURRENT_SCHEMA_VERSION,
));
}
let response_type_raw = header.get_u8();
let response_id = ResponseId::from_u128(header.get_u128_le());
let handler_name_len = header.get_u16_le() as usize;
if handler_name_len == 0 || header.remaining() < handler_name_len + 2 {
return Err(DecodeError::InvalidHandlerNameLength);
}
let handler_name_bytes = header.copy_to_bytes(handler_name_len);
let handler_name =
String::from_utf8(handler_name_bytes.to_vec()).map_err(|_| DecodeError::InvalidUtf8)?;
let response_type = ResponseType::try_from(response_type_raw)?;
let headers_len = header.get_u16_le() as usize;
let headers = if headers_len > 0 {
if headers_len > MAX_HEADERS_LEN || header.remaining() < headers_len {
return Err(DecodeError::InvalidHeadersLength);
}
let headers_bytes = header.copy_to_bytes(headers_len);
let headers_map: HashMap<String, String> = rmp_serde::from_slice(&headers_bytes)?;
Some(headers_map)
} else {
None
};
Ok(ActiveMessage {
metadata: MessageMetadata {
schema_version,
response_type,
response_id,
handler_name,
headers,
},
payload,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decode_response_id_from_request_header_roundtrip() {
let response_id = ResponseId::from_u128(0xDEAD_BEEF_CAFE_F00D_1234_5678_90AB_CDEF);
let (header, _, _) = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, "h".to_string(), None),
payload: Bytes::from_static(b""),
}
.encode()
.unwrap();
let decoded = decode_response_id_from_request_header(&header).unwrap();
assert_eq!(decoded.as_u128(), response_id.as_u128());
}
#[test]
fn decode_response_id_rejects_truncated_header() {
let short = Bytes::from_static(&[1u8, 0, 0, 0]);
assert!(decode_response_id_from_request_header(&short).is_none());
}
#[test]
fn decode_response_id_rejects_wrong_schema() {
let mut bad = vec![0u8; 18];
bad[1] = 2; let header = Bytes::from(bad);
assert!(decode_response_id_from_request_header(&header).is_none());
}
#[test]
fn decode_response_id_rejects_invalid_response_type() {
let mut bad = vec![0u8; 18];
bad[0] = CURRENT_SCHEMA_VERSION;
bad[1] = 99; let header = Bytes::from(bad);
assert!(decode_response_id_from_request_header(&header).is_none());
}
#[test]
fn decode_response_id_rejects_response_format_header() {
let header = Bytes::from(vec![1u8, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]);
let _ = decode_response_id_from_request_header(&header);
}
#[test]
fn test_handler_name_at_u16_max_succeeds() {
let handler_name = "a".repeat(u16::MAX as usize);
let response_id = ResponseId::from_u128(12345);
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), None),
payload: Bytes::from_static(b"test payload"),
};
let result = message.encode();
assert!(
result.is_ok(),
"Handler name at u16::MAX should encode successfully"
);
let (header, payload, _) = result.unwrap();
let decoded = decode_active_message(header, payload).unwrap();
assert_eq!(decoded.metadata.handler_name, handler_name);
}
#[test]
fn test_handler_name_exceeds_u16_max_fails() {
let handler_name = "a".repeat(u16::MAX as usize + 1);
let response_id = ResponseId::from_u128(12345);
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), None),
payload: Bytes::from_static(b"test payload"),
};
let result = message.encode();
assert!(
result.is_err(),
"Handler name exceeding u16::MAX should fail to encode"
);
match result {
Err(EncodeError::HandlerNameTooLong(len)) => {
assert_eq!(len, u16::MAX as usize + 1);
}
_ => panic!("Expected HandlerNameTooLong error"),
}
}
#[test]
fn test_handler_name_way_too_long_fails() {
let handler_name = "a".repeat(1024 * 1024);
let response_id = ResponseId::from_u128(12345);
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name, None),
payload: Bytes::from_static(b"test payload"),
};
let result = message.encode();
assert!(
result.is_err(),
"Very large handler name should fail to encode"
);
}
#[test]
fn test_normal_handler_name_succeeds() {
let handler_name = "my_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), None),
payload: Bytes::from_static(b"test payload"),
};
let (header, payload, _) = message.encode().unwrap();
let decoded = decode_active_message(header, payload).unwrap();
assert_eq!(decoded.metadata.handler_name, handler_name);
}
#[test]
fn test_headers_encode_decode_round_trip() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
headers.insert("trace-id".to_string(), "abc123".to_string());
headers.insert("span-id".to_string(), "def456".to_string());
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(
response_id,
handler_name.clone(),
Some(headers.clone()),
),
payload: Bytes::from_static(b"test payload"),
};
let (header, payload, _) = message.encode().unwrap();
let decoded = decode_active_message(header, payload).unwrap();
assert_eq!(decoded.metadata.handler_name, handler_name);
assert_eq!(
decoded.metadata.response_id.as_u128(),
response_id.as_u128()
);
assert!(decoded.metadata.headers.is_some());
let decoded_headers = decoded.metadata.headers.unwrap();
assert_eq!(decoded_headers.len(), 2);
assert_eq!(decoded_headers.get("trace-id").unwrap(), "abc123");
assert_eq!(decoded_headers.get("span-id").unwrap(), "def456");
}
#[test]
fn test_headers_none_encodes_with_zero_length() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), None),
payload: Bytes::from_static(b"test payload"),
};
let (header, payload, _) = message.encode().unwrap();
let expected_len = 1 + 1 + 16 + 2 + handler_name.len() + 2;
assert_eq!(header.len(), expected_len);
let decoded = decode_active_message(header, payload).unwrap();
assert!(decoded.metadata.headers.is_none());
}
#[test]
fn test_headers_empty_map_encodes_successfully() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let headers = HashMap::new();
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), Some(headers)),
payload: Bytes::from_static(b"test payload"),
};
let (header, payload, _) = message.encode().unwrap();
let decoded = decode_active_message(header, payload).unwrap();
assert!(decoded.metadata.headers.is_some());
assert_eq!(decoded.metadata.headers.unwrap().len(), 0);
}
#[test]
fn test_headers_per_value_size_limit() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
let value_1kb = "a".repeat(1024);
headers.insert("large-header".to_string(), value_1kb);
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), Some(headers)),
payload: Bytes::from_static(b"test payload"),
};
let result = message.encode();
assert!(result.is_ok(), "1KB value should encode successfully");
}
#[test]
fn test_headers_per_value_size_exceeds_limit() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
let value_too_large = "a".repeat(1025);
headers.insert("large-header".to_string(), value_too_large);
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), Some(headers)),
payload: Bytes::from_static(b"test payload"),
};
let result = message.encode();
assert!(result.is_err(), "1KB+1 value should fail to encode");
match result {
Err(EncodeError::HeaderValueTooLarge(key, size)) => {
assert_eq!(key, "large-header");
assert_eq!(size, 1025);
}
_ => panic!("Expected HeaderValueTooLarge error"),
}
}
#[test]
fn test_headers_total_size_limit() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
for i in 0..40 {
let key = format!("header-{}", i);
let value = "x".repeat(500);
headers.insert(key, value);
}
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, handler_name.clone(), Some(headers)),
payload: Bytes::from_static(b"test payload"),
};
let result = message.encode();
assert!(result.is_err(), "Total size exceeding 16KB should fail");
match result {
Err(EncodeError::TotalHeadersTooLarge(size)) => {
assert!(size > 16384, "Size should exceed 16KB");
}
_ => panic!("Expected TotalHeadersTooLarge error"),
}
}
#[test]
fn test_headers_with_special_characters() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
headers.insert("emoji".to_string(), "ππ".to_string());
headers.insert("unicode".to_string(), "δ½ ε₯½δΈη".to_string());
headers.insert("special".to_string(), "a\nb\tc\"d'e".to_string());
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(
response_id,
handler_name.clone(),
Some(headers.clone()),
),
payload: Bytes::from_static(b"test payload"),
};
let (header, payload, _) = message.encode().unwrap();
let decoded = decode_active_message(header, payload).unwrap();
let decoded_headers = decoded.metadata.headers.unwrap();
assert_eq!(decoded_headers.get("emoji").unwrap(), "ππ");
assert_eq!(decoded_headers.get("unicode").unwrap(), "δ½ ε₯½δΈη");
assert_eq!(decoded_headers.get("special").unwrap(), "a\nb\tc\"d'e");
}
#[test]
fn test_headers_with_many_entries() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
for i in 0..100 {
headers.insert(format!("key-{}", i), format!("value-{}", i));
}
let message = ActiveMessage {
metadata: MessageMetadata::new_unary(
response_id,
handler_name.clone(),
Some(headers.clone()),
),
payload: Bytes::from_static(b"test payload"),
};
let (header, payload, _) = message.encode().unwrap();
let decoded = decode_active_message(header, payload).unwrap();
let decoded_headers = decoded.metadata.headers.unwrap();
assert_eq!(decoded_headers.len(), 100);
assert_eq!(decoded_headers.get("key-42").unwrap(), "value-42");
}
#[test]
fn test_headers_all_response_types() {
let handler_name = "test_handler".to_string();
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
headers.insert("test".to_string(), "value".to_string());
let msg_fire = ActiveMessage {
metadata: MessageMetadata::new_fire(
response_id,
handler_name.clone(),
Some(headers.clone()),
),
payload: Bytes::from_static(b"test"),
};
let (h, p, _) = msg_fire.encode().unwrap();
let decoded = decode_active_message(h, p).unwrap();
assert_eq!(decoded.metadata.response_type, ResponseType::FireAndForget);
assert!(decoded.metadata.headers.is_some());
let msg_sync = ActiveMessage {
metadata: MessageMetadata::new_sync(
response_id,
handler_name.clone(),
Some(headers.clone()),
),
payload: Bytes::from_static(b"test"),
};
let (h, p, _) = msg_sync.encode().unwrap();
let decoded = decode_active_message(h, p).unwrap();
assert_eq!(decoded.metadata.response_type, ResponseType::AckNack);
assert!(decoded.metadata.headers.is_some());
let msg_unary = ActiveMessage {
metadata: MessageMetadata::new_unary(
response_id,
handler_name.clone(),
Some(headers.clone()),
),
payload: Bytes::from_static(b"test"),
};
let (h, p, _) = msg_unary.encode().unwrap();
let decoded = decode_active_message(h, p).unwrap();
assert_eq!(decoded.metadata.response_type, ResponseType::Unary);
assert!(decoded.metadata.headers.is_some());
}
}