diff --git a/core/connectors/runtime/src/sink.rs b/core/connectors/runtime/src/sink.rs index 7a17724510..02d1ce7671 100644 --- a/core/connectors/runtime/src/sink.rs +++ b/core/connectors/runtime/src/sink.rs @@ -602,9 +602,10 @@ async fn process_messages( } let decode_elapsed = decode_start.elapsed(); - let mut messages = Vec::with_capacity(decoded.len()); + let mut batches = Vec::::new(); for message in decoded { let mut current_message = Some(message); + let mut transform_failed = false; for transform in transforms.iter() { let Some(message) = current_message.take() else { break; @@ -621,12 +622,16 @@ async fn process_messages( topic_metadata.topic ); error_count += 1; - current_message = None; + transform_failed = true; break; } } } + if transform_failed { + continue; + } + // Filter contract: transform returning Ok(None) is an intentional drop. let Some(message) = current_message else { filtered_count += 1; @@ -673,6 +678,7 @@ async fn process_messages( continue; }; + let payload_schema = message.payload.schema(); let Ok(payload) = message.payload.try_into_vec() else { error!( "Failed to get message payload for message with ID: {id}, offset: {offset} for sink connector with ID: {plugin_id}" @@ -695,7 +701,7 @@ async fn process_messages( None => vec![], }; - messages.push(RawMessage { + let raw_message = RawMessage { id, offset, checksum, @@ -703,7 +709,14 @@ async fn process_messages( origin_timestamp, headers, payload, - }); + }; + match batches.last_mut() { + Some(batch) if batch.schema == payload_schema => batch.messages.push(raw_message), + _ => batches.push(RawMessages { + schema: payload_schema, + messages: vec![raw_message], + }), + } } metrics.inc_errors_by_with_labels(&labels.counter, error_count); @@ -711,7 +724,13 @@ async fn process_messages( metrics.inc_messages_filtered_with_labels(&labels.counter, filtered_count); } - let processed_count = messages.len(); + let processed_count = batches.iter().map(|batch| batch.messages.len()).sum(); + if batches.is_empty() { + batches.push(RawMessages { + schema: messages_metadata.schema, + messages: Vec::new(), + }); + } let topic_meta = postcard::to_allocvec(topic_metadata).map_err(|error| { error!( @@ -720,33 +739,36 @@ async fn process_messages( RuntimeError::FailedToSerializeTopicMetadata })?; - let messages_meta = postcard::to_allocvec(&messages_metadata).map_err(|error| { - error!( - "Failed to serialize messages metadata for sink connector with ID: {plugin_id}. {error}" - ); - RuntimeError::FailedToSerializeMessagesMetadata - })?; - - let messages = postcard::to_allocvec(&RawMessages { - schema: decoder.schema(), - messages, - }) - .map_err(|error| { - error!("Failed to serialize messages for sink connector with ID: {plugin_id}. {error}"); - RuntimeError::FailedToSerializeRawMessages - })?; + let mut ffi_elapsed = Duration::ZERO; + for batch in batches { + let messages_metadata = MessagesMetadata { + partition_id: messages_metadata.partition_id, + current_offset: messages_metadata.current_offset, + schema: batch.schema, + }; + let messages_meta = postcard::to_allocvec(&messages_metadata).map_err(|error| { + error!( + "Failed to serialize messages metadata for sink connector with ID: {plugin_id}. {error}" + ); + RuntimeError::FailedToSerializeMessagesMetadata + })?; + let messages = postcard::to_allocvec(&batch).map_err(|error| { + error!("Failed to serialize messages for sink connector with ID: {plugin_id}. {error}"); + RuntimeError::FailedToSerializeRawMessages + })?; - let ffi_start = Instant::now(); - (consume)( - plugin_id, - topic_meta.as_ptr(), - topic_meta.len(), - messages_meta.as_ptr(), - messages_meta.len(), - messages.as_ptr(), - messages.len(), - ); - let ffi_elapsed = ffi_start.elapsed(); + let ffi_start = Instant::now(); + (consume)( + plugin_id, + topic_meta.as_ptr(), + topic_meta.len(), + messages_meta.as_ptr(), + messages_meta.len(), + messages.as_ptr(), + messages.len(), + ); + ffi_elapsed += ffi_start.elapsed(); + } Ok(SinkBatchTiming { processed_count, @@ -760,3 +782,470 @@ struct SinkBatchTiming { decode_elapsed: Duration, ffi_elapsed: Duration, } + +#[cfg(test)] +mod tests { + use std::{path::PathBuf, sync::Mutex, thread}; + + use async_trait::async_trait; + use iggy_connector_sdk::{ + ConsumedMessage, Error, Payload, Sink, + sink::SinkContainer, + transforms::{ProtoConvert, ProtoConvertConfig, TransformType}, + }; + + use super::*; + + static CAPTURED_BATCHES: Mutex> = Mutex::new(Vec::new()); + static CONSUMED_BATCHES: Mutex> = Mutex::new(Vec::new()); + static SINK_CONTAINER: Mutex>> = Mutex::new(None); + + struct CapturedBatch { + metadata_schema: Schema, + batch_schema: Schema, + offsets: Vec, + payloads: Vec>, + } + + struct ConsumedBatch { + metadata_schema: Schema, + offsets: Vec, + } + + #[derive(Debug)] + struct RecordingSink { + plugin_id: u32, + } + + #[async_trait] + impl Sink for RecordingSink { + async fn open(&mut self) -> Result<(), Error> { + Ok(()) + } + + async fn consume( + &self, + _topic_metadata: &TopicMetadata, + messages_metadata: MessagesMetadata, + messages: Vec, + ) -> Result<(), Error> { + CONSUMED_BATCHES + .lock() + .expect("consumed batches lock should succeed") + .push(( + self.plugin_id, + ConsumedBatch { + metadata_schema: messages_metadata.schema, + offsets: messages.into_iter().map(|message| message.offset).collect(), + }, + )); + Ok(()) + } + + async fn close(&mut self) -> Result<(), Error> { + Ok(()) + } + } + + struct RawToText; + + impl Transform for RawToText { + fn r#type(&self) -> TransformType { + TransformType::ProtoConvert + } + + fn transform( + &self, + _metadata: &TopicMetadata, + message: DecodedMessage, + ) -> Result, Error> { + let Payload::Raw(payload) = message.payload else { + return Ok(Some(message)); + }; + let text = String::from_utf8(payload).map_err(|_| Error::InvalidTextPayload)?; + Ok(Some(DecodedMessage { + payload: Payload::Text(text), + ..message + })) + } + } + + struct FailingTransform; + + impl Transform for FailingTransform { + fn r#type(&self) -> TransformType { + TransformType::ProtoConvert + } + + fn transform( + &self, + _metadata: &TopicMetadata, + _message: DecodedMessage, + ) -> Result, Error> { + Err(Error::InvalidPayloadType) + } + } + + extern "C" fn capture_schemas( + plugin_id: u32, + _topic_meta_ptr: *const u8, + _topic_meta_len: usize, + messages_meta_ptr: *const u8, + messages_meta_len: usize, + messages_ptr: *const u8, + messages_len: usize, + ) -> i32 { + unsafe { + let messages_metadata = postcard::from_bytes::( + std::slice::from_raw_parts(messages_meta_ptr, messages_meta_len), + ) + .expect("messages metadata should deserialize"); + let raw_messages = postcard::from_bytes::(std::slice::from_raw_parts( + messages_ptr, + messages_len, + )) + .expect("raw messages should deserialize"); + CAPTURED_BATCHES + .lock() + .expect("capture lock should succeed") + .push(( + plugin_id, + CapturedBatch { + metadata_schema: messages_metadata.schema, + batch_schema: raw_messages.schema, + offsets: raw_messages + .messages + .iter() + .map(|message| message.offset) + .collect(), + payloads: raw_messages + .messages + .into_iter() + .map(|message| message.payload) + .collect(), + }, + )); + } + 0 + } + + extern "C" fn consume_through_sdk( + _plugin_id: u32, + topic_meta_ptr: *const u8, + topic_meta_len: usize, + messages_meta_ptr: *const u8, + messages_meta_len: usize, + messages_ptr: *const u8, + messages_len: usize, + ) -> i32 { + let (topic_metadata, messages_metadata, messages) = unsafe { + ( + std::slice::from_raw_parts(topic_meta_ptr, topic_meta_len).to_vec(), + std::slice::from_raw_parts(messages_meta_ptr, messages_meta_len).to_vec(), + std::slice::from_raw_parts(messages_ptr, messages_len).to_vec(), + ) + }; + thread::spawn(move || { + let container = SINK_CONTAINER + .lock() + .expect("sink container lock should succeed"); + let Some(container) = container.as_ref() else { + return -1; + }; + unsafe { + container.consume( + topic_metadata.as_ptr(), + topic_metadata.len(), + messages_metadata.as_ptr(), + messages_metadata.len(), + messages.as_ptr(), + messages.len(), + ) + } + }) + .join() + .expect("SDK consume thread should succeed") + } + + extern "C" fn discard_log( + _level: u8, + _target_ptr: *const u8, + _target_len: usize, + _message_ptr: *const u8, + _message_len: usize, + ) { + } + + fn take_captured_batches(plugin_id: u32) -> Vec { + let mut captured = CAPTURED_BATCHES + .lock() + .expect("capture lock should succeed"); + let all_batches = std::mem::take(&mut *captured); + let (matching, remaining): (Vec<_>, Vec<_>) = all_batches + .into_iter() + .partition(|(captured_plugin_id, _)| *captured_plugin_id == plugin_id); + *captured = remaining; + matching.into_iter().map(|(_, batch)| batch).collect() + } + + fn take_consumed_batches(plugin_id: u32) -> Vec { + let mut consumed = CONSUMED_BATCHES + .lock() + .expect("consumed batches lock should succeed"); + let all_batches = std::mem::take(&mut *consumed); + let (matching, remaining): (Vec<_>, Vec<_>) = all_batches + .into_iter() + .partition(|(consumed_plugin_id, _)| *consumed_plugin_id == plugin_id); + *consumed = remaining; + matching.into_iter().map(|(_, batch)| batch).collect() + } + + fn proto_convert() -> Arc { + let schema_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("../sdk/examples"); + let schema_path = schema_dir.join("user.proto"); + Arc::new(ProtoConvert::new(ProtoConvertConfig { + source_format: Schema::Json, + target_format: Schema::Proto, + schema_path: Some(schema_path), + message_type: Some("com.example.User".to_string()), + include_paths: vec![schema_dir], + ..ProtoConvertConfig::default() + })) + } + + fn mixed_schema_messages() -> Vec { + vec![ + message(1, 1, "[]"), + message(2, 2, r#"{"id":2,"name":"second"}"#), + message(3, 3, r#"{"id":3,"name":"third"}"#), + message(4, 4, "null"), + message(5, 5, r#"{"id":5,"name":"fifth"}"#), + ] + } + + fn message(id: u128, offset: u64, payload: &'static str) -> IggyMessage { + let mut message = IggyMessage::builder() + .id(id) + .payload(payload.into()) + .build() + .expect("message should build"); + message.header.offset = offset; + message + } + + #[tokio::test] + async fn given_transform_changes_schema_when_crossing_ffi_should_use_output_schema() { + let metrics = Arc::new(Metrics::init()); + let labels = SinkLabels::new("schema-aware"); + let decoder: Arc = Schema::Raw.decoder(); + let transforms: Vec> = vec![Arc::new(RawToText)]; + let consume: ConsumeCallback = capture_schemas; + let plugin_id = 1; + let message = message(1, 0, "transformed"); + + let timing = process_messages( + plugin_id, + MessagesMetadata { + partition_id: 1, + current_offset: 0, + schema: Schema::Raw, + }, + &TopicMetadata { + stream: "stream".to_string(), + topic: "topic".to_string(), + }, + vec![message], + &consume, + &transforms, + &decoder, + &metrics, + &labels, + ) + .await + .expect("message processing should succeed"); + + let mut captured = take_captured_batches(plugin_id); + assert_eq!(captured.len(), 1); + let captured = captured.pop().expect("FFI callback should capture schemas"); + assert_eq!(timing.processed_count, 1); + assert_eq!(captured.metadata_schema, Schema::Text); + assert_eq!(captured.batch_schema, Schema::Text); + assert_eq!(captured.offsets, vec![0]); + assert_eq!(captured.payloads, vec![b"transformed".to_vec()]); + } + + #[tokio::test] + async fn given_mixed_transform_schemas_when_crossing_ffi_should_preserve_all_messages_in_order() + { + let plugin_id = 2; + let metrics = Arc::new(Metrics::init()); + let labels = SinkLabels::new("mixed-schema"); + let decoder: Arc = Schema::Json.decoder(); + let transforms = vec![proto_convert()]; + let consume: ConsumeCallback = capture_schemas; + + let timing = process_messages( + plugin_id, + MessagesMetadata { + partition_id: 1, + current_offset: 5, + schema: Schema::Json, + }, + &TopicMetadata { + stream: "stream".to_string(), + topic: "topic".to_string(), + }, + mixed_schema_messages(), + &consume, + &transforms, + &decoder, + &metrics, + &labels, + ) + .await + .expect("message processing should succeed"); + + let captured = take_captured_batches(plugin_id); + assert_eq!(timing.processed_count, 5); + assert_eq!( + metrics.get_errors("mixed-schema", crate::metrics::ConnectorType::Sink), + 0 + ); + assert_eq!(captured.len(), 4); + assert_eq!( + captured + .iter() + .map(|batch| batch.metadata_schema) + .collect::>(), + vec![Schema::Proto, Schema::Raw, Schema::Proto, Schema::Raw] + ); + assert!( + captured + .iter() + .all(|batch| batch.metadata_schema == batch.batch_schema) + ); + assert_eq!( + captured + .iter() + .flat_map(|batch| batch.offsets.iter().copied()) + .collect::>(), + vec![1, 2, 3, 4, 5] + ); + } + + #[tokio::test] + async fn given_transform_failure_when_processing_should_count_error_without_filtering() { + let plugin_id = 3; + let connector_key = "transform-failure"; + let metrics = Arc::new(Metrics::init()); + let labels = SinkLabels::new(connector_key); + let decoder: Arc = Schema::Raw.decoder(); + let transforms: Vec> = vec![Arc::new(FailingTransform)]; + let consume: ConsumeCallback = capture_schemas; + + let timing = process_messages( + plugin_id, + MessagesMetadata { + partition_id: 1, + current_offset: 1, + schema: Schema::Raw, + }, + &TopicMetadata { + stream: "stream".to_string(), + topic: "topic".to_string(), + }, + vec![message(1, 1, "invalid")], + &consume, + &transforms, + &decoder, + &metrics, + &labels, + ) + .await + .expect("message processing should succeed"); + + assert_eq!(timing.processed_count, 0); + assert_eq!( + metrics.get_errors(connector_key, crate::metrics::ConnectorType::Sink), + 1 + ); + assert_eq!( + metrics.get_messages_filtered(connector_key, crate::metrics::ConnectorType::Sink), + 0 + ); + take_captured_batches(plugin_id); + } + + #[test] + fn given_mixed_transform_schemas_when_consumed_by_sdk_should_reconstruct_every_message() { + let plugin_id = 4; + let metrics = Arc::new(Metrics::init()); + let labels = SinkLabels::new("sdk-schema-reconstruction"); + let decoder: Arc = Schema::Json.decoder(); + let transforms = vec![proto_convert()]; + let consume: ConsumeCallback = consume_through_sdk; + let config = b"{}"; + let mut container = SinkContainer::new(plugin_id); + let open_result = unsafe { + container.open::<_, serde_json::Value>( + plugin_id, + config.as_ptr(), + config.len(), + discard_log, + |plugin_id, _config| RecordingSink { plugin_id }, + ) + }; + assert_eq!(open_result, 0); + *SINK_CONTAINER + .lock() + .expect("sink container lock should succeed") = Some(container); + + let runtime = tokio::runtime::Runtime::new().expect("test runtime should initialize"); + let timing = runtime + .block_on(process_messages( + plugin_id, + MessagesMetadata { + partition_id: 1, + current_offset: 5, + schema: Schema::Json, + }, + &TopicMetadata { + stream: "stream".to_string(), + topic: "topic".to_string(), + }, + mixed_schema_messages(), + &consume, + &transforms, + &decoder, + &metrics, + &labels, + )) + .expect("message processing should succeed"); + drop(runtime); + + let consumed = take_consumed_batches(plugin_id); + assert_eq!(timing.processed_count, 5); + assert_eq!(consumed.len(), 4); + assert_eq!( + consumed + .iter() + .map(|batch| batch.metadata_schema) + .collect::>(), + vec![Schema::Proto, Schema::Raw, Schema::Proto, Schema::Raw] + ); + assert_eq!( + consumed + .iter() + .flat_map(|batch| batch.offsets.iter().copied()) + .collect::>(), + vec![1, 2, 3, 4, 5] + ); + + let mut container = SINK_CONTAINER + .lock() + .expect("sink container lock should succeed") + .take() + .expect("sink container should be initialized"); + assert_eq!(unsafe { container.close() }, 0); + } +} diff --git a/core/connectors/sdk/src/lib.rs b/core/connectors/sdk/src/lib.rs index c8ed2ff94d..90f2446cea 100644 --- a/core/connectors/sdk/src/lib.rs +++ b/core/connectors/sdk/src/lib.rs @@ -145,6 +145,17 @@ pub enum Payload { } impl Payload { + pub const fn schema(&self) -> Schema { + match self { + Self::Json(_) => Schema::Json, + Self::Raw(_) => Schema::Raw, + Self::Text(_) => Schema::Text, + Self::Proto(_) => Schema::Proto, + Self::FlatBuffer(_) => Schema::FlatBuffer, + Self::Avro(_) => Schema::Avro, + } + } + /// Consuming conversion — transfers ownership of inner buffers. pub fn try_into_vec(self) -> Result, Error> { match self {