use crate::Identifier;
use crate::IggyMessageView;
use crate::PartitioningKind;
use crate::Validatable;
use crate::error::IggyError;
use crate::types::message::HeaderEntry;
use crate::types::message::partitioning::Partitioning;
use crate::{IggyMessage, IggyMessagesBatch};
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
use bytes::Bytes;
use serde::de::{self, MapAccess, Visitor};
use serde::ser::SerializeStruct;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::collections::BTreeMap;
use std::fmt::Formatter;
#[derive(Debug, PartialEq)]
pub struct SendMessages {
pub metadata_length: u32,
pub stream_id: Identifier,
pub topic_id: Identifier,
pub partitioning: Partitioning,
pub batch: IggyMessagesBatch,
}
impl Default for SendMessages {
fn default() -> Self {
SendMessages {
metadata_length: 0,
stream_id: Identifier::default(),
topic_id: Identifier::default(),
partitioning: Partitioning::default(),
batch: IggyMessagesBatch::empty(),
}
}
}
impl Validatable<IggyError> for SendMessages {
fn validate(&self) -> Result<(), IggyError> {
if self.partitioning.value.len() > 255
|| (self.partitioning.kind != PartitioningKind::Balanced
&& self.partitioning.value.is_empty())
{
return Err(IggyError::InvalidKeyValueLength);
}
self.batch.validate()?;
Ok(())
}
}
impl Serialize for SendMessages {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let messages: Vec<serde_json::Value> = self
.batch
.iter()
.map(|msg_view: IggyMessageView<'_>| {
let mut obj = serde_json::json!({
"id": msg_view.header().id(),
"payload": BASE64.encode(msg_view.payload()),
});
match msg_view.user_headers_map() {
Ok(Some(headers)) => {
let entries: Vec<HeaderEntry> = headers
.into_iter()
.map(|(k, v)| HeaderEntry { key: k, value: v })
.collect();
obj["user_headers"] = serde_json::to_value(&entries).unwrap();
}
_ if msg_view.user_headers().is_some() => {
let raw_base64 = BASE64.encode(msg_view.user_headers().unwrap());
obj["user_headers"] = serde_json::to_value(raw_base64).unwrap();
}
_ => {}
}
obj
})
.collect();
let mut state = serializer.serialize_struct("SendMessages", 2)?;
state.serialize_field("partitioning", &self.partitioning)?;
state.serialize_field("messages", &messages)?;
state.end()
}
}
impl<'de> Deserialize<'de> for SendMessages {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
enum Field {
Partitioning,
Messages,
}
impl<'de> Deserialize<'de> for Field {
fn deserialize<D>(deserializer: D) -> Result<Field, D::Error>
where
D: Deserializer<'de>,
{
struct FieldVisitor;
impl Visitor<'_> for FieldVisitor {
type Value = Field;
fn expecting(&self, formatter: &mut Formatter) -> std::fmt::Result {
formatter.write_str("`partitioning` or `messages`")
}
fn visit_str<E>(self, value: &str) -> Result<Field, E>
where
E: de::Error,
{
match value {
"partitioning" => Ok(Field::Partitioning),
"messages" => Ok(Field::Messages),
_ => Err(de::Error::unknown_field(
value,
&["partitioning", "messages"],
)),
}
}
}
deserializer.deserialize_identifier(FieldVisitor)
}
}
struct SendMessagesVisitor;
impl<'de> Visitor<'de> for SendMessagesVisitor {
type Value = SendMessages;
fn expecting(&self, formatter: &mut Formatter) -> std::fmt::Result {
formatter.write_str("struct SendMessages")
}
fn visit_map<V>(self, mut map: V) -> Result<SendMessages, V::Error>
where
V: MapAccess<'de>,
{
let mut partitioning = None;
let mut messages = None;
while let Some(key) = map.next_key()? {
match key {
Field::Partitioning => {
if partitioning.is_some() {
return Err(de::Error::duplicate_field("partitioning"));
}
partitioning = Some(map.next_value()?);
}
Field::Messages => {
if messages.is_some() {
return Err(de::Error::duplicate_field("messages"));
}
let message_data: Vec<serde_json::Value> = map.next_value()?;
let mut iggy_messages = Vec::new();
for msg in message_data {
let id = parse_message_id(msg.get("id")).map_err(|error| {
de::Error::custom(format!("Invalid message ID: {error}"))
})?;
let payload = msg
.get("payload")
.and_then(|v| v.as_str())
.ok_or_else(|| de::Error::missing_field("payload"))?;
let payload_bytes = BASE64
.decode(payload)
.map_err(|_| de::Error::custom("Invalid base64 payload"))?;
let (headers_map, raw_headers) =
if let Some(headers) = msg.get("user_headers") {
if headers.is_null() {
(None, None)
} else if let Some(base64_str) = headers.as_str() {
let raw = BASE64.decode(base64_str).map_err(|e| {
de::Error::custom(format!(
"Invalid base64 headers: {e}"
))
})?;
(None, Some(Bytes::from(raw)))
} else {
let entries: Vec<HeaderEntry> = serde_json::from_value(
headers.clone(),
)
.map_err(|e| {
de::Error::custom(format!(
"Invalid headers format: {e}"
))
})?;
let mut map = BTreeMap::new();
for entry in entries {
map.insert(entry.key, entry.value);
}
(Some(map), None)
}
} else {
(None, None)
};
let mut iggy_message = if let Some(headers) = headers_map {
IggyMessage::builder()
.id(id)
.payload(payload_bytes.into())
.user_headers(headers)
.build()
.map_err(|e| {
de::Error::custom(format!(
"Failed to create message with headers: {e}"
))
})?
} else {
IggyMessage::builder()
.id(id)
.payload(payload_bytes.into())
.build()
.map_err(|e| {
de::Error::custom(format!(
"Failed to create message: {e}"
))
})?
};
if let Some(raw) = raw_headers {
iggy_message.header.user_headers_length = raw.len() as u32;
iggy_message.user_headers = Some(raw);
}
iggy_messages.push(iggy_message);
}
messages = Some(iggy_messages);
}
}
}
let partitioning =
partitioning.ok_or_else(|| de::Error::missing_field("partitioning"))?;
let messages = messages.ok_or_else(|| de::Error::missing_field("messages"))?;
let batch = IggyMessagesBatch::from(&messages);
Ok(SendMessages {
metadata_length: 0, stream_id: Identifier::default(),
topic_id: Identifier::default(),
partitioning,
batch,
})
}
}
deserializer.deserialize_struct(
"SendMessages",
&["partitioning", "messages"],
SendMessagesVisitor,
)
}
}
fn parse_message_id(value: Option<&serde_json::Value>) -> Result<u128, String> {
let value = match value {
Some(v) => v,
None => return Ok(0),
};
match value {
serde_json::Value::Number(id) => id
.as_u64()
.map(|v| v as u128)
.ok_or_else(|| "ID must be a positive integer".to_string()),
serde_json::Value::String(id) => {
if let Ok(id) = id.parse::<u128>() {
return Ok(id);
}
let hex_str = id.replace('-', "");
if hex_str.len() == 32 && hex_str.chars().all(|c| c.is_ascii_hexdigit()) {
u128::from_str_radix(&hex_str, 16)
.map_err(|error| format!("Invalid UUID format: {error}"))
} else {
Err(format!(
"Invalid ID string: '{id}' - must be a decimal number or UUID hex format",
))
}
}
_ => Err("ID must be a number or string".to_string()),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deserialize_send_messages_with_invalid_uuid_fails() {
let json_data = serde_json::json!({
"partitioning": {
"kind": "balanced",
"value": ""
},
"messages": [{
"id": "114514-invalid-uuid-1919810",
"payload": "SGVsbG8gSWdneSE=",
"user_headers": [{
"key": "content-type",
"value": "text/plain"
}]
}]
});
assert!(serde_json::from_value::<SendMessages>(json_data).is_err());
}
#[test]
fn key_of_type_balanced_should_have_empty_value() {
let key = Partitioning::balanced();
assert_eq!(key.kind, PartitioningKind::Balanced);
assert_eq!(key.length, 0);
assert!(key.value.is_empty());
assert_eq!(
PartitioningKind::from_code(1).unwrap(),
PartitioningKind::Balanced
);
}
#[test]
fn key_of_type_partition_should_have_value_of_const_length_4() {
let partition_id = 1234u32;
let key = Partitioning::partition_id(partition_id);
assert_eq!(key.kind, PartitioningKind::PartitionId);
assert_eq!(key.length, 4);
assert_eq!(key.value, partition_id.to_le_bytes());
assert_eq!(
PartitioningKind::from_code(2).unwrap(),
PartitioningKind::PartitionId
);
}
#[test]
fn key_of_type_messages_key_should_have_value_of_dynamic_length() {
let messages_key = "hello world";
let key = Partitioning::messages_key_str(messages_key).unwrap();
assert_eq!(key.kind, PartitioningKind::MessagesKey);
assert_eq!(key.length, messages_key.len() as u8);
assert_eq!(key.value, messages_key.as_bytes());
assert_eq!(
PartitioningKind::from_code(3).unwrap(),
PartitioningKind::MessagesKey
);
}
#[test]
fn key_of_type_messages_key_that_has_length_0_should_fail() {
let messages_key = "";
let key = Partitioning::messages_key_str(messages_key);
assert!(key.is_err());
}
#[test]
fn key_of_type_messages_key_that_has_length_greater_than_255_should_fail() {
let messages_key = "a".repeat(256);
let key = Partitioning::messages_key_str(&messages_key);
assert!(key.is_err());
}
#[test]
fn parse_message_id_from_number() {
let value = serde_json::json!(12345);
let id = parse_message_id(Some(&value)).unwrap();
assert_eq!(id, 12345u128);
}
#[test]
fn parse_message_id_from_large_number_string() {
let value = serde_json::json!("340282366920938463463374607431768211455");
let id = parse_message_id(Some(&value)).unwrap();
assert_eq!(id, 340282366920938463463374607431768211455u128);
}
#[test]
fn parse_message_id_from_uuid_with_dashes() {
let value = serde_json::json!("af362865-042c-4000-0000-000000000000");
let id = parse_message_id(Some(&value)).unwrap();
assert_eq!(id, 0xaf362865042c40000000000000000000u128);
}
#[test]
fn parse_message_id_from_uuid_without_dashes() {
let value = serde_json::json!("af362865042c40000000000000000000");
let id = parse_message_id(Some(&value)).unwrap();
assert_eq!(id, 0xaf362865042c40000000000000000000u128);
}
#[test]
fn parse_message_id_defaults_to_zero_when_missing() {
let id = parse_message_id(None).unwrap();
assert_eq!(id, 0u128);
}
#[test]
fn parse_message_id_rejects_invalid_string() {
let value = serde_json::json!("not-a-valid-id");
let result = parse_message_id(Some(&value));
assert!(result.is_err());
}
}