From 8557f0fd81573adce8a8a580af2de71c611009b2 Mon Sep 17 00:00:00 2001 From: haubur Date: Sat, 22 Aug 2026 14:17:49 +0200 Subject: [PATCH 1/2] fix(sdk): error callback should await shard flush on producer shutdown --- core/sdk/src/clients/producer_dispatcher.rs | 78 +++++++++++++++------ 1 file changed, 58 insertions(+), 20 deletions(-) diff --git a/core/sdk/src/clients/producer_dispatcher.rs b/core/sdk/src/clients/producer_dispatcher.rs index 1cfe1f7515..e036ec64dd 100644 --- a/core/sdk/src/clients/producer_dispatcher.rs +++ b/core/sdk/src/clients/producer_dispatcher.rs @@ -49,29 +49,17 @@ impl ProducerDispatcher { let err_callback = config.error_callback.clone(); let (stop_tx, _) = broadcast::channel::<()>(1); - let mut stop_rx = stop_tx.subscribe(); + // A task that receives errors from shards when writing messages failed. let handle = tokio::spawn(async move { - loop { - tokio::select! { - maybe_message = err_rx.recv_async() => { - match maybe_message { - Ok(ctx) => { - if let Err(panic) = std::panic::AssertUnwindSafe(err_callback.call(ctx)) - .catch_unwind() - .await - { - tracing::error!("error_callback panicked: {:?}", panic); - } - } - Err(_) => break - } - } - _ = stop_rx.recv() => { - tracing::debug!("error-callback worker finished"); - break - } + while let Ok(ctx) = err_rx.recv_async().await { + if let Err(panic) = std::panic::AssertUnwindSafe(err_callback.call(ctx)) + .catch_unwind() + .await + { + tracing::error!("error_callback panicked: {:?}", panic); } } + tracing::debug!("error-callback worker finished"); }); let bytes_permit = { @@ -192,6 +180,8 @@ impl ProducerDispatcher { } } + // After shards are closed await the error callback task, + // that might drain queued errors from the final flush. if let Err(e) = self._join_handle.await { tracing::error!("error-worker panicked: {e:?}"); } @@ -432,4 +422,52 @@ mod tests { assert_eq!(sharding_called.load(Ordering::SeqCst), 1); assert_eq!(error_called.load(Ordering::SeqCst), 1); } + + #[tokio::test] + async fn test_shutdown_reports_errors_from_final_flush() { + let mut mock = MockProducerCoreBackend::new(); + // mock a failed send to be caught by the error callback + mock.expect_send_internal() + .times(1) + .returning(|_, _, _, _| { + Box::pin(async { + Err(IggyError::ProducerSendFailed { + cause: Box::new(IggyError::Error), + failed: Arc::new(vec![dummy_message(10)]), + committed: Arc::new(Vec::new()), + stream_name: "1".to_string(), + topic_name: "1".to_string(), + }) + }) + }); + + let error_called = Arc::new(AtomicUsize::new(0)); + + let config = BackgroundConfig::builder() + .num_shards(1) + .linger_time(Duration::from_secs(60).into()) // block/ wait so shutdown triggers the flush + .error_callback(Arc::new(Box::new(TestErrorCallback { + called: error_called.clone(), + }))) + .build(); + + let dispatcher = ProducerDispatcher::new(Arc::new(mock), config); + + dispatcher + .dispatch( + vec![dummy_message(10)], + dummy_identifier(), + dummy_identifier(), + None, + ) + .await + .unwrap(); + + // trigger the final flush + dispatcher.shutdown().await; + + // Passes, if the error callback is hit once after final + // flush from shutdown. + assert_eq!(error_called.load(Ordering::SeqCst), 1); + } } From f555aa7ab0623a550e6f0bb6297cd72e78795dcf Mon Sep 17 00:00:00 2001 From: haubur Date: Tue, 25 Aug 2026 17:00:09 +0200 Subject: [PATCH 2/2] explicitly assert on ErrorCallback ctx.messages --- core/sdk/src/clients/producer_dispatcher.rs | 68 ++++++++++----------- 1 file changed, 34 insertions(+), 34 deletions(-) diff --git a/core/sdk/src/clients/producer_dispatcher.rs b/core/sdk/src/clients/producer_dispatcher.rs index e036ec64dd..acb79c45bf 100644 --- a/core/sdk/src/clients/producer_dispatcher.rs +++ b/core/sdk/src/clients/producer_dispatcher.rs @@ -32,7 +32,7 @@ pub struct ProducerDispatcher { closed: AtomicBool, bytes_permit: Arc, stop_tx: broadcast::Sender<()>, - _join_handle: JoinHandle<()>, + join_handle: JoinHandle<()>, } impl ProducerDispatcher { @@ -49,7 +49,6 @@ impl ProducerDispatcher { let err_callback = config.error_callback.clone(); let (stop_tx, _) = broadcast::channel::<()>(1); - // A task that receives errors from shards when writing messages failed. let handle = tokio::spawn(async move { while let Ok(ctx) = err_rx.recv_async().await { if let Err(panic) = std::panic::AssertUnwindSafe(err_callback.call(ctx)) @@ -90,7 +89,7 @@ impl ProducerDispatcher { closed: AtomicBool::new(false), bytes_permit: Arc::new(Semaphore::new(bytes_permit)), stop_tx, - _join_handle: handle, + join_handle: handle, } } @@ -182,7 +181,7 @@ impl ProducerDispatcher { // After shards are closed await the error callback task, // that might drain queued errors from the final flush. - if let Err(e) = self._join_handle.await { + if let Err(e) = self.join_handle.await { tracing::error!("error-worker panicked: {e:?}"); } } @@ -366,11 +365,14 @@ mod tests { #[derive(Clone, Debug)] struct TestErrorCallback { called: Arc, + last_batch_len: Arc, } impl ErrorCallback for TestErrorCallback { - fn call(&self, _ctx: ErrorCtx) -> Pin + Send + 'static>> { + fn call(&self, ctx: ErrorCtx) -> Pin + Send + 'static>> { self.called.fetch_add(1, Ordering::SeqCst); + self.last_batch_len + .store(ctx.messages.len(), Ordering::SeqCst); Box::pin(async {}) } } @@ -378,27 +380,27 @@ mod tests { #[tokio::test] async fn test_custom_sharding_and_error_callback() { let mut mock = MockProducerCoreBackend::new(); - mock.expect_send_internal() - .times(1) - .returning(|_, _, _, _| { - Box::pin(async { - Err(IggyError::ProducerSendFailed { - cause: Box::new(IggyError::Error), - failed: Arc::new(vec![dummy_message(10)]), - committed: Arc::new(Vec::new()), - stream_name: "1".to_string(), - topic_name: "1".to_string(), - }) + mock.expect_send_internal().returning(|_, _, _, _| { + Box::pin(async { + Err(IggyError::ProducerSendFailed { + cause: Box::new(IggyError::Error), + failed: Arc::new(vec![dummy_message(10)]), + committed: Arc::new(Vec::new()), + stream_name: "1".to_string(), + topic_name: "1".to_string(), }) - }); + }) + }); let sharding_called = Arc::new(AtomicUsize::new(0)); let error_called = Arc::new(AtomicUsize::new(0)); + let last_batch_len = Arc::new(AtomicUsize::new(0)); let config = BackgroundConfig::builder() .num_shards(1) .error_callback(Arc::new(Box::new(TestErrorCallback { called: error_called.clone(), + last_batch_len: last_batch_len.clone(), }))) .sharding(Box::new(TestSharding { called: sharding_called.clone(), @@ -421,33 +423,33 @@ mod tests { assert!(result.is_ok()); assert_eq!(sharding_called.load(Ordering::SeqCst), 1); assert_eq!(error_called.load(Ordering::SeqCst), 1); + assert_eq!(last_batch_len.load(Ordering::SeqCst), 1); } #[tokio::test] async fn test_shutdown_reports_errors_from_final_flush() { let mut mock = MockProducerCoreBackend::new(); - // mock a failed send to be caught by the error callback - mock.expect_send_internal() - .times(1) - .returning(|_, _, _, _| { - Box::pin(async { - Err(IggyError::ProducerSendFailed { - cause: Box::new(IggyError::Error), - failed: Arc::new(vec![dummy_message(10)]), - committed: Arc::new(Vec::new()), - stream_name: "1".to_string(), - topic_name: "1".to_string(), - }) + mock.expect_send_internal().returning(|_, _, _, _| { + Box::pin(async { + Err(IggyError::ProducerSendFailed { + cause: Box::new(IggyError::Error), + failed: Arc::new(vec![dummy_message(10)]), + committed: Arc::new(Vec::new()), + stream_name: "1".to_string(), + topic_name: "1".to_string(), }) - }); + }) + }); let error_called = Arc::new(AtomicUsize::new(0)); + let last_batch_len = Arc::new(AtomicUsize::new(0)); let config = BackgroundConfig::builder() .num_shards(1) - .linger_time(Duration::from_secs(60).into()) // block/ wait so shutdown triggers the flush + .linger_time(Duration::from_secs(60).into()) .error_callback(Arc::new(Box::new(TestErrorCallback { called: error_called.clone(), + last_batch_len: last_batch_len.clone(), }))) .build(); @@ -463,11 +465,9 @@ mod tests { .await .unwrap(); - // trigger the final flush dispatcher.shutdown().await; - // Passes, if the error callback is hit once after final - // flush from shutdown. assert_eq!(error_called.load(Ordering::SeqCst), 1); + assert_eq!(last_batch_len.load(Ordering::SeqCst), 1); } }