use ahash::AHashMap as HashMap;
use bytes::Bytes;
use std::any::{Any, TypeId};
use std::fmt;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::sync::Arc;
use crate::consumer::ConsumerRecord;
use crate::error::KrafkaError;
use crate::producer::{ProducerRecord, RecordHeaders, RecordMetadata};
use crate::{Offset, PartitionId, Timestamp};
pub type CommitOffsets = HashMap<(String, PartitionId), Offset>;
pub type InterceptorResult = std::result::Result<(), Box<dyn std::error::Error + Send + Sync>>;
struct ContextEntry {
owner: u16,
type_id: TypeId,
value: Box<dyn Any + Send + Sync>,
}
#[derive(Default)]
pub struct RecordContext {
entries: Vec<ContextEntry>,
owner: u16,
}
impl fmt::Debug for RecordContext {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RecordContext")
.field("entries", &self.entries.len())
.field("owner", &self.owner)
.finish()
}
}
impl RecordContext {
#[must_use]
pub const fn new() -> Self {
Self {
entries: Vec::new(),
owner: 0,
}
}
pub(crate) fn set_owner(&mut self, owner: u16) -> u16 {
std::mem::replace(&mut self.owner, owner)
}
fn position<T: 'static>(&self) -> Option<usize> {
let type_id = TypeId::of::<T>();
self.entries
.iter()
.position(|entry| entry.owner == self.owner && entry.type_id == type_id)
}
pub fn insert<T: Send + Sync + 'static>(&mut self, value: T) -> Option<T> {
match self.position::<T>() {
Some(index) => {
let previous = self
.entries
.get_mut(index)
.map(|entry| std::mem::replace(&mut entry.value, Box::new(value)))?;
previous.downcast::<T>().ok().map(|boxed| *boxed)
}
None => {
self.entries.push(ContextEntry {
owner: self.owner,
type_id: TypeId::of::<T>(),
value: Box::new(value),
});
None
}
}
}
#[must_use]
pub fn get<T: Send + Sync + 'static>(&self) -> Option<&T> {
let index = self.position::<T>()?;
self.entries.get(index)?.value.downcast_ref::<T>()
}
#[must_use]
pub fn get_mut<T: Send + Sync + 'static>(&mut self) -> Option<&mut T> {
let index = self.position::<T>()?;
self.entries.get_mut(index)?.value.downcast_mut::<T>()
}
#[must_use]
pub fn take<T: Send + Sync + 'static>(&mut self) -> Option<T> {
let index = self.position::<T>()?;
let entry = self.entries.remove(index);
entry.value.downcast::<T>().ok().map(|boxed| *boxed)
}
#[must_use]
pub fn contains<T: Send + Sync + 'static>(&self) -> bool {
self.position::<T>().is_some()
}
}
pub trait ProducerInterceptor: Send + Sync + fmt::Debug {
fn on_send(&self, _record: &mut ProducerRecord, _ctx: &mut RecordContext) -> InterceptorResult {
Ok(())
}
fn on_acknowledgement(
&self,
_metadata: &RecordMetadata,
_error: Option<&KrafkaError>,
_headers: &RecordHeaders,
_ctx: &mut RecordContext,
) -> InterceptorResult {
Ok(())
}
fn close(&self) -> InterceptorResult {
Ok(())
}
}
pub trait ConsumerInterceptor: Send + Sync + fmt::Debug {
fn on_consume(&self, _records: &[ConsumerRecord]) -> InterceptorResult {
Ok(())
}
fn on_commit(
&self,
_offsets: &CommitOffsets,
_error: Option<&KrafkaError>,
) -> InterceptorResult {
Ok(())
}
fn close(&self) -> InterceptorResult {
Ok(())
}
}
#[derive(Debug)]
pub(crate) struct NoOpProducerInterceptor;
impl ProducerInterceptor for NoOpProducerInterceptor {}
#[derive(Debug)]
pub(crate) struct NoOpConsumerInterceptor;
impl ConsumerInterceptor for NoOpConsumerInterceptor {}
struct CheapRecordSnapshot {
partition: Option<PartitionId>,
key: Option<Bytes>,
value: Option<Bytes>,
timestamp: Option<Timestamp>,
header_len: usize,
}
impl CheapRecordSnapshot {
#[inline]
fn capture(record: &ProducerRecord) -> Self {
Self {
partition: record.partition,
key: record.key.clone(),
value: record.value.clone(),
timestamp: record.timestamp,
header_len: record.headers.len(),
}
}
#[inline]
fn restore(self, record: &mut ProducerRecord) {
record.partition = self.partition;
record.key = self.key;
record.value = self.value;
record.timestamp = self.timestamp;
record.headers.truncate(self.header_len);
}
}
pub(crate) struct ProducerInterceptorChain {
interceptors: Vec<Arc<dyn ProducerInterceptor>>,
}
#[inline]
fn chain_owner(index: usize) -> u16 {
u16::try_from(index).unwrap_or(u16::MAX)
}
impl fmt::Debug for ProducerInterceptorChain {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ProducerInterceptorChain")
.field("len", &self.interceptors.len())
.finish()
}
}
impl ProducerInterceptorChain {
pub fn new(interceptors: Vec<Arc<dyn ProducerInterceptor>>) -> Self {
Self { interceptors }
}
}
impl ProducerInterceptor for ProducerInterceptorChain {
fn on_send(&self, record: &mut ProducerRecord, ctx: &mut RecordContext) -> InterceptorResult {
for (i, interceptor) in self.interceptors.iter().enumerate() {
let previous_owner = ctx.set_owner(chain_owner(i));
let snapshot = CheapRecordSnapshot::capture(record);
match catch_unwind(AssertUnwindSafe(|| interceptor.on_send(record, ctx))) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
chain_index = i,
chain_len = self.interceptors.len(),
topic = record.topic.as_str(),
error = %e,
"ProducerInterceptor.on_send failed",
);
}
Err(_) => {
snapshot.restore(record);
tracing::error!(
chain_index = i,
chain_len = self.interceptors.len(),
topic = record.topic.as_str(),
"ProducerInterceptor.on_send panicked — record partially restored (payload redacted)",
);
}
}
ctx.set_owner(previous_owner);
}
Ok(())
}
fn on_acknowledgement(
&self,
metadata: &RecordMetadata,
error: Option<&KrafkaError>,
headers: &RecordHeaders,
ctx: &mut RecordContext,
) -> InterceptorResult {
for (i, interceptor) in self.interceptors.iter().enumerate() {
let previous_owner = ctx.set_owner(chain_owner(i));
let outcome = catch_unwind(AssertUnwindSafe(|| {
interceptor.on_acknowledgement(metadata, error, headers, ctx)
}));
ctx.set_owner(previous_owner);
match outcome {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
chain_index = i,
chain_len = self.interceptors.len(),
topic = metadata.topic.as_str(),
partition = metadata.partition,
error = %e,
"ProducerInterceptor.on_acknowledgement failed",
);
}
Err(_) => {
tracing::error!(
chain_index = i,
chain_len = self.interceptors.len(),
topic = metadata.topic.as_str(),
partition = metadata.partition,
"ProducerInterceptor.on_acknowledgement panicked (payload redacted)",
);
}
}
}
Ok(())
}
fn close(&self) -> InterceptorResult {
for (i, interceptor) in self.interceptors.iter().enumerate() {
match catch_unwind(AssertUnwindSafe(|| interceptor.close())) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
chain_index = i,
chain_len = self.interceptors.len(),
error = %e,
"ProducerInterceptor.close failed",
);
}
Err(_) => {
tracing::error!(
chain_index = i,
chain_len = self.interceptors.len(),
"ProducerInterceptor.close panicked (payload redacted)",
);
}
}
}
Ok(())
}
}
pub(crate) struct ConsumerInterceptorChain {
interceptors: Vec<Arc<dyn ConsumerInterceptor>>,
}
impl fmt::Debug for ConsumerInterceptorChain {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConsumerInterceptorChain")
.field("len", &self.interceptors.len())
.finish()
}
}
impl ConsumerInterceptorChain {
pub fn new(interceptors: Vec<Arc<dyn ConsumerInterceptor>>) -> Self {
Self { interceptors }
}
}
impl ConsumerInterceptor for ConsumerInterceptorChain {
fn on_consume(&self, records: &[ConsumerRecord]) -> InterceptorResult {
for (i, interceptor) in self.interceptors.iter().enumerate() {
match catch_unwind(AssertUnwindSafe(|| interceptor.on_consume(records))) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
chain_index = i,
chain_len = self.interceptors.len(),
record_count = records.len(),
error = %e,
"ConsumerInterceptor.on_consume failed",
);
}
Err(_) => {
tracing::error!(
chain_index = i,
chain_len = self.interceptors.len(),
record_count = records.len(),
"ConsumerInterceptor.on_consume panicked (payload redacted)",
);
}
}
}
Ok(())
}
fn on_commit(
&self,
offsets: &HashMap<(String, PartitionId), Offset>,
error: Option<&KrafkaError>,
) -> InterceptorResult {
for (i, interceptor) in self.interceptors.iter().enumerate() {
match catch_unwind(AssertUnwindSafe(|| interceptor.on_commit(offsets, error))) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
chain_index = i,
chain_len = self.interceptors.len(),
offset_count = offsets.len(),
error = %e,
"ConsumerInterceptor.on_commit failed",
);
}
Err(_) => {
tracing::error!(
chain_index = i,
chain_len = self.interceptors.len(),
offset_count = offsets.len(),
"ConsumerInterceptor.on_commit panicked (payload redacted)",
);
}
}
}
Ok(())
}
fn close(&self) -> InterceptorResult {
for (i, interceptor) in self.interceptors.iter().enumerate() {
match catch_unwind(AssertUnwindSafe(|| interceptor.close())) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
chain_index = i,
chain_len = self.interceptors.len(),
error = %e,
"ConsumerInterceptor.close failed",
);
}
Err(_) => {
tracing::error!(
chain_index = i,
chain_len = self.interceptors.len(),
"ConsumerInterceptor.close panicked (payload redacted)",
);
}
}
}
Ok(())
}
}
pub(crate) fn safe_on_send(
interceptor: &dyn ProducerInterceptor,
record: &mut ProducerRecord,
ctx: &mut RecordContext,
) {
match catch_unwind(AssertUnwindSafe(|| interceptor.on_send(record, ctx))) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
topic = record.topic.as_str(),
error = %e,
"ProducerInterceptor.on_send failed",
);
}
Err(_) => {
tracing::error!(
topic = record.topic.as_str(),
"ProducerInterceptor.on_send panicked (payload redacted)",
);
}
}
}
pub(crate) fn safe_on_acknowledgement(
interceptor: &dyn ProducerInterceptor,
metadata: &RecordMetadata,
error: Option<&KrafkaError>,
headers: &RecordHeaders,
ctx: &mut RecordContext,
) {
match catch_unwind(AssertUnwindSafe(|| {
interceptor.on_acknowledgement(metadata, error, headers, ctx)
})) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
topic = metadata.topic.as_str(),
partition = metadata.partition,
error = %e,
"ProducerInterceptor.on_acknowledgement failed",
);
}
Err(_) => {
tracing::error!(
topic = metadata.topic.as_str(),
partition = metadata.partition,
"ProducerInterceptor.on_acknowledgement panicked (payload redacted)",
);
}
}
}
pub(crate) fn safe_producer_close(interceptor: &dyn ProducerInterceptor) {
match catch_unwind(AssertUnwindSafe(|| interceptor.close())) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
error = %e,
"ProducerInterceptor.close failed",
);
}
Err(_) => {
tracing::error!("ProducerInterceptor.close panicked (payload redacted)");
}
}
}
pub(crate) fn safe_on_consume(interceptor: &dyn ConsumerInterceptor, records: &[ConsumerRecord]) {
match catch_unwind(AssertUnwindSafe(|| interceptor.on_consume(records))) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
record_count = records.len(),
error = %e,
"ConsumerInterceptor.on_consume failed",
);
}
Err(_) => {
tracing::error!(
record_count = records.len(),
"ConsumerInterceptor.on_consume panicked (payload redacted)",
);
}
}
}
pub(crate) fn safe_on_commit(
interceptor: &dyn ConsumerInterceptor,
offsets: &HashMap<(String, PartitionId), Offset>,
error: Option<&KrafkaError>,
) {
match catch_unwind(AssertUnwindSafe(|| interceptor.on_commit(offsets, error))) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
offset_count = offsets.len(),
error = %e,
"ConsumerInterceptor.on_commit failed",
);
}
Err(_) => {
tracing::error!(
offset_count = offsets.len(),
"ConsumerInterceptor.on_commit panicked (payload redacted)",
);
}
}
}
pub(crate) fn safe_consumer_close(interceptor: &dyn ConsumerInterceptor) {
match catch_unwind(AssertUnwindSafe(|| interceptor.close())) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!(
error = %e,
"ConsumerInterceptor.close failed",
);
}
Err(_) => {
tracing::error!("ConsumerInterceptor.close panicked (payload redacted)");
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
#[derive(Debug)]
struct TestProducerInterceptor {
send_count: std::sync::atomic::AtomicUsize,
ack_count: std::sync::atomic::AtomicUsize,
}
impl TestProducerInterceptor {
fn new() -> Self {
Self {
send_count: std::sync::atomic::AtomicUsize::new(0),
ack_count: std::sync::atomic::AtomicUsize::new(0),
}
}
fn send_count(&self) -> usize {
self.send_count.load(std::sync::atomic::Ordering::Relaxed)
}
fn ack_count(&self) -> usize {
self.ack_count.load(std::sync::atomic::Ordering::Relaxed)
}
}
impl ProducerInterceptor for TestProducerInterceptor {
fn on_send(
&self,
record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
self.send_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
record.headers.push((
"x-intercepted".to_string(),
Some(bytes::Bytes::from_static(b"true")),
));
Ok(())
}
fn on_acknowledgement(
&self,
_metadata: &RecordMetadata,
_error: Option<&KrafkaError>,
_headers: &RecordHeaders,
_ctx: &mut RecordContext,
) -> InterceptorResult {
self.ack_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(())
}
}
#[derive(Debug)]
struct TestConsumerInterceptor {
consume_count: std::sync::atomic::AtomicUsize,
commit_count: std::sync::atomic::AtomicUsize,
}
impl TestConsumerInterceptor {
fn new() -> Self {
Self {
consume_count: std::sync::atomic::AtomicUsize::new(0),
commit_count: std::sync::atomic::AtomicUsize::new(0),
}
}
fn consume_count(&self) -> usize {
self.consume_count
.load(std::sync::atomic::Ordering::Relaxed)
}
fn commit_count(&self) -> usize {
self.commit_count.load(std::sync::atomic::Ordering::Relaxed)
}
}
impl ConsumerInterceptor for TestConsumerInterceptor {
fn on_consume(&self, records: &[ConsumerRecord]) -> InterceptorResult {
self.consume_count
.fetch_add(records.len(), std::sync::atomic::Ordering::Relaxed);
Ok(())
}
fn on_commit(
&self,
_offsets: &HashMap<(String, PartitionId), Offset>,
_error: Option<&KrafkaError>,
) -> InterceptorResult {
self.commit_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(())
}
}
#[test]
fn test_producer_interceptor_on_send() {
let interceptor = TestProducerInterceptor::new();
let mut record = ProducerRecord::new("test-topic", b"value".to_vec());
assert_eq!(interceptor.send_count(), 0);
assert!(record.headers.is_empty());
interceptor
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
assert_eq!(interceptor.send_count(), 1);
assert_eq!(record.headers.len(), 1);
assert_eq!(record.headers[0].0, "x-intercepted");
assert_eq!(
record.headers[0].1,
Some(bytes::Bytes::from_static(b"true"))
);
}
#[test]
fn test_producer_interceptor_on_acknowledgement() {
let interceptor = TestProducerInterceptor::new();
let metadata = RecordMetadata {
topic: "test-topic".to_string(),
partition: 0,
offset: 42,
timestamp: 1000,
delivery: crate::producer::DeliveryConfirmation::Offset,
};
interceptor
.on_acknowledgement(&metadata, None, &[], &mut RecordContext::new())
.unwrap();
assert_eq!(interceptor.ack_count(), 1);
let err = KrafkaError::config("test error");
interceptor
.on_acknowledgement(&metadata, Some(&err), &[], &mut RecordContext::new())
.unwrap();
assert_eq!(interceptor.ack_count(), 2);
}
#[test]
fn test_consumer_interceptor_on_consume() {
let interceptor = TestConsumerInterceptor::new();
let records = vec![
ConsumerRecord::new("test-topic", 0, 0, None, Some(bytes::Bytes::from("v1"))),
ConsumerRecord::new("test-topic", 0, 1, None, Some(bytes::Bytes::from("v2"))),
];
interceptor.on_consume(&records).unwrap();
assert_eq!(interceptor.consume_count(), 2);
}
#[test]
fn test_consumer_interceptor_on_commit() {
let interceptor = TestConsumerInterceptor::new();
let mut offsets = HashMap::new();
offsets.insert(("test-topic".to_string(), 0), 10i64);
interceptor.on_commit(&offsets, None).unwrap();
assert_eq!(interceptor.commit_count(), 1);
}
#[test]
fn test_noop_interceptors() {
let producer_interceptor = NoOpProducerInterceptor;
let mut record = ProducerRecord::new("test", b"value".to_vec());
producer_interceptor
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
assert!(record.headers.is_empty());
let consumer_interceptor = NoOpConsumerInterceptor;
consumer_interceptor.on_consume(&[]).unwrap();
consumer_interceptor
.on_commit(&HashMap::new(), None)
.unwrap();
}
#[derive(Debug)]
struct PanickingProducerInterceptor;
impl ProducerInterceptor for PanickingProducerInterceptor {
fn on_send(
&self,
_record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
panic!("on_send panic");
}
fn on_acknowledgement(
&self,
_metadata: &RecordMetadata,
_error: Option<&KrafkaError>,
_headers: &RecordHeaders,
_ctx: &mut RecordContext,
) -> InterceptorResult {
panic!("on_acknowledgement panic");
}
fn close(&self) -> InterceptorResult {
panic!("producer close panic");
}
}
#[derive(Debug)]
struct PanickingConsumerInterceptor;
impl ConsumerInterceptor for PanickingConsumerInterceptor {
fn on_consume(&self, _records: &[ConsumerRecord]) -> InterceptorResult {
panic!("on_consume panic");
}
fn on_commit(
&self,
_offsets: &HashMap<(String, PartitionId), Offset>,
_error: Option<&KrafkaError>,
) -> InterceptorResult {
panic!("on_commit panic");
}
fn close(&self) -> InterceptorResult {
panic!("consumer close panic");
}
}
#[test]
fn test_safe_on_send_catches_panic() {
let interceptor = PanickingProducerInterceptor;
let mut record = ProducerRecord::new("test", b"value".to_vec());
safe_on_send(&interceptor, &mut record, &mut RecordContext::new());
}
#[test]
fn test_safe_on_acknowledgement_catches_panic() {
let interceptor = PanickingProducerInterceptor;
let metadata = RecordMetadata {
topic: "test".to_string(),
partition: 0,
offset: 0,
timestamp: 0,
delivery: crate::producer::DeliveryConfirmation::Offset,
};
safe_on_acknowledgement(
&interceptor,
&metadata,
None,
&[],
&mut RecordContext::new(),
);
}
#[test]
fn test_safe_producer_close_catches_panic() {
let interceptor = PanickingProducerInterceptor;
safe_producer_close(&interceptor);
}
#[test]
fn test_safe_on_consume_catches_panic() {
let interceptor = PanickingConsumerInterceptor;
safe_on_consume(&interceptor, &[]);
}
#[test]
fn test_safe_on_commit_catches_panic() {
let interceptor = PanickingConsumerInterceptor;
safe_on_commit(&interceptor, &HashMap::new(), None);
}
#[test]
fn test_safe_consumer_close_catches_panic() {
let interceptor = PanickingConsumerInterceptor;
safe_consumer_close(&interceptor);
}
#[test]
fn test_close_default_noop() {
let p = NoOpProducerInterceptor;
p.close().unwrap();
let c = NoOpConsumerInterceptor;
c.close().unwrap();
}
#[derive(Debug)]
struct OrderedProducerInterceptor {
name: &'static str,
log: Arc<std::sync::Mutex<Vec<String>>>,
}
impl ProducerInterceptor for OrderedProducerInterceptor {
fn on_send(
&self,
_record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
self.log
.lock()
.unwrap()
.push(format!("{}.on_send", self.name));
Ok(())
}
fn on_acknowledgement(
&self,
_metadata: &RecordMetadata,
_error: Option<&KrafkaError>,
_headers: &RecordHeaders,
_ctx: &mut RecordContext,
) -> InterceptorResult {
self.log
.lock()
.unwrap()
.push(format!("{}.on_ack", self.name));
Ok(())
}
fn close(&self) -> InterceptorResult {
self.log
.lock()
.unwrap()
.push(format!("{}.close", self.name));
Ok(())
}
}
#[test]
fn test_producer_chain_executes_in_order() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ProducerInterceptorChain::new(vec![
Arc::new(OrderedProducerInterceptor {
name: "first",
log: Arc::clone(&log),
}),
Arc::new(OrderedProducerInterceptor {
name: "second",
log: Arc::clone(&log),
}),
Arc::new(OrderedProducerInterceptor {
name: "third",
log: Arc::clone(&log),
}),
]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
let metadata = RecordMetadata {
topic: "test".to_string(),
partition: 0,
offset: 0,
timestamp: 0,
delivery: crate::producer::DeliveryConfirmation::Offset,
};
chain
.on_acknowledgement(&metadata, None, &[], &mut RecordContext::new())
.unwrap();
chain.close().unwrap();
let log = log.lock().unwrap();
assert_eq!(
*log,
vec![
"first.on_send",
"second.on_send",
"third.on_send",
"first.on_ack",
"second.on_ack",
"third.on_ack",
"first.close",
"second.close",
"third.close",
]
);
}
#[test]
fn test_producer_chain_on_send_mutations_visible_to_next() {
#[derive(Debug)]
struct HeaderAdder(&'static str);
impl ProducerInterceptor for HeaderAdder {
fn on_send(
&self,
record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
record.headers.push((
self.0.to_string(),
Some(bytes::Bytes::copy_from_slice(self.0.as_bytes())),
));
Ok(())
}
}
let chain = ProducerInterceptorChain::new(vec![
Arc::new(HeaderAdder("first")),
Arc::new(HeaderAdder("second")),
]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
assert_eq!(record.headers.len(), 2);
assert_eq!(record.headers[0].0, "first");
assert_eq!(record.headers[1].0, "second");
}
#[test]
fn test_producer_chain_panic_isolation() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ProducerInterceptorChain::new(vec![
Arc::new(OrderedProducerInterceptor {
name: "before",
log: Arc::clone(&log),
}),
Arc::new(PanickingProducerInterceptor),
Arc::new(OrderedProducerInterceptor {
name: "after",
log: Arc::clone(&log),
}),
]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
let metadata = RecordMetadata {
topic: "test".to_string(),
partition: 0,
offset: 0,
timestamp: 0,
delivery: crate::producer::DeliveryConfirmation::Offset,
};
chain
.on_acknowledgement(&metadata, None, &[], &mut RecordContext::new())
.unwrap();
chain.close().unwrap();
let log = log.lock().unwrap();
assert_eq!(
*log,
vec![
"before.on_send",
"after.on_send",
"before.on_ack",
"after.on_ack",
"before.close",
"after.close",
]
);
}
#[derive(Debug)]
struct MutateThenPanicInterceptor;
impl ProducerInterceptor for MutateThenPanicInterceptor {
fn on_send(
&self,
record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
record.partition = Some(99);
record.timestamp = Some(1234);
record.key = Some(bytes::Bytes::from_static(b"clobbered-key"));
record.value = Some(bytes::Bytes::from_static(b"clobbered-value"));
record
.headers
.push(("added-before-panic".to_string(), Some(bytes::Bytes::new())));
record.topic = "clobbered-topic".to_string();
panic!("mutate then panic");
}
}
#[test]
fn test_producer_chain_panic_restores_cheap_fields() {
let chain = ProducerInterceptorChain::new(vec![Arc::new(MutateThenPanicInterceptor)]);
let mut record = ProducerRecord::new("original-topic", b"original-value".to_vec());
record.key = Some(bytes::Bytes::from_static(b"original-key"));
record.partition = Some(1);
record.timestamp = Some(7);
record.headers.push((
"pre-existing".to_string(),
Some(bytes::Bytes::from_static(b"h")),
));
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
assert_eq!(record.partition, Some(1));
assert_eq!(record.timestamp, Some(7));
assert_eq!(record.key, Some(bytes::Bytes::from_static(b"original-key")));
assert_eq!(record.value, Some(bytes::Bytes::from("original-value")));
assert_eq!(record.headers.len(), 1);
assert_eq!(record.headers[0].0, "pre-existing");
}
#[test]
fn test_producer_chain_panic_does_not_deep_clone_topic() {
let chain = ProducerInterceptorChain::new(vec![Arc::new(MutateThenPanicInterceptor)]);
let mut record = ProducerRecord::new("original-topic", b"v".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
assert_eq!(record.topic, "clobbered-topic");
}
#[test]
fn test_producer_chain_panic_still_surfaced_and_chain_continues() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
#[derive(Debug)]
struct Observer(Arc<std::sync::Mutex<Vec<String>>>);
impl ProducerInterceptor for Observer {
fn on_send(
&self,
record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
self.0.lock().unwrap().push(format!(
"value={} headers={}",
record.value_str().unwrap_or("<null>"),
record.headers.len()
));
Ok(())
}
}
let chain = ProducerInterceptorChain::new(vec![
Arc::new(MutateThenPanicInterceptor),
Arc::new(Observer(Arc::clone(&log))),
]);
let mut record = ProducerRecord::new("t", b"v".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
let log = log.lock().unwrap();
assert_eq!(*log, vec!["value=v headers=0"]);
}
#[test]
fn test_producer_chain_no_panic_keeps_mutations() {
#[derive(Debug)]
struct Mutator;
impl ProducerInterceptor for Mutator {
fn on_send(
&self,
record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
record.partition = Some(5);
record.value = Some(bytes::Bytes::from_static(b"new"));
record
.headers
.push(("added".to_string(), Some(bytes::Bytes::new())));
Ok(())
}
}
let chain = ProducerInterceptorChain::new(vec![Arc::new(Mutator)]);
let mut record = ProducerRecord::new("t", b"old".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
assert_eq!(record.partition, Some(5));
assert_eq!(record.value, Some(bytes::Bytes::from_static(b"new")));
assert_eq!(record.headers.len(), 1);
}
#[test]
fn test_producer_chain_error_return_does_not_roll_back() {
#[derive(Debug)]
struct MutateThenErr;
impl ProducerInterceptor for MutateThenErr {
fn on_send(
&self,
record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
record.partition = Some(3);
Err("boom".into())
}
}
let chain = ProducerInterceptorChain::new(vec![Arc::new(MutateThenErr)]);
let mut record = ProducerRecord::new("t", b"v".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
assert_eq!(record.partition, Some(3));
}
#[test]
fn test_producer_chain_empty() {
let chain = ProducerInterceptorChain::new(vec![]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
chain.close().unwrap();
}
#[derive(Debug)]
struct OrderedConsumerInterceptor {
name: &'static str,
log: Arc<std::sync::Mutex<Vec<String>>>,
}
impl ConsumerInterceptor for OrderedConsumerInterceptor {
fn on_consume(&self, _records: &[ConsumerRecord]) -> InterceptorResult {
self.log
.lock()
.unwrap()
.push(format!("{}.on_consume", self.name));
Ok(())
}
fn on_commit(
&self,
_offsets: &HashMap<(String, PartitionId), Offset>,
_error: Option<&KrafkaError>,
) -> InterceptorResult {
self.log
.lock()
.unwrap()
.push(format!("{}.on_commit", self.name));
Ok(())
}
fn close(&self) -> InterceptorResult {
self.log
.lock()
.unwrap()
.push(format!("{}.close", self.name));
Ok(())
}
}
#[test]
fn test_consumer_chain_executes_in_order() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ConsumerInterceptorChain::new(vec![
Arc::new(OrderedConsumerInterceptor {
name: "first",
log: Arc::clone(&log),
}),
Arc::new(OrderedConsumerInterceptor {
name: "second",
log: Arc::clone(&log),
}),
]);
chain.on_consume(&[]).unwrap();
chain.on_commit(&HashMap::new(), None).unwrap();
chain.close().unwrap();
let log = log.lock().unwrap();
assert_eq!(
*log,
vec![
"first.on_consume",
"second.on_consume",
"first.on_commit",
"second.on_commit",
"first.close",
"second.close",
]
);
}
#[test]
fn test_consumer_chain_panic_isolation() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ConsumerInterceptorChain::new(vec![
Arc::new(OrderedConsumerInterceptor {
name: "before",
log: Arc::clone(&log),
}),
Arc::new(PanickingConsumerInterceptor),
Arc::new(OrderedConsumerInterceptor {
name: "after",
log: Arc::clone(&log),
}),
]);
chain.on_consume(&[]).unwrap();
chain.on_commit(&HashMap::new(), None).unwrap();
chain.close().unwrap();
let log = log.lock().unwrap();
assert_eq!(
*log,
vec![
"before.on_consume",
"after.on_consume",
"before.on_commit",
"after.on_commit",
"before.close",
"after.close",
]
);
}
#[test]
fn test_consumer_chain_empty() {
let chain = ConsumerInterceptorChain::new(vec![]);
chain.on_consume(&[]).unwrap();
chain.on_commit(&HashMap::new(), None).unwrap();
chain.close().unwrap();
}
#[test]
fn test_chain_via_safe_wrappers() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ProducerInterceptorChain::new(vec![
Arc::new(OrderedProducerInterceptor {
name: "a",
log: Arc::clone(&log),
}),
Arc::new(OrderedProducerInterceptor {
name: "b",
log: Arc::clone(&log),
}),
]);
let mut record = ProducerRecord::new("test", b"v".to_vec());
safe_on_send(&chain, &mut record, &mut RecordContext::new());
let log = log.lock().unwrap();
assert_eq!(*log, vec!["a.on_send", "b.on_send"]);
}
#[derive(Debug)]
struct FailingProducerInterceptor;
impl ProducerInterceptor for FailingProducerInterceptor {
fn on_send(
&self,
_record: &mut ProducerRecord,
_ctx: &mut RecordContext,
) -> InterceptorResult {
Err("metrics backend unavailable".into())
}
fn on_acknowledgement(
&self,
_metadata: &RecordMetadata,
_error: Option<&KrafkaError>,
_headers: &RecordHeaders,
_ctx: &mut RecordContext,
) -> InterceptorResult {
Err("ack handler failed".into())
}
fn close(&self) -> InterceptorResult {
Err("cleanup failed".into())
}
}
#[derive(Debug)]
struct FailingConsumerInterceptor;
impl ConsumerInterceptor for FailingConsumerInterceptor {
fn on_consume(&self, _records: &[ConsumerRecord]) -> InterceptorResult {
Err("consume handler failed".into())
}
fn on_commit(
&self,
_offsets: &HashMap<(String, PartitionId), Offset>,
_error: Option<&KrafkaError>,
) -> InterceptorResult {
Err("commit handler failed".into())
}
fn close(&self) -> InterceptorResult {
Err("cleanup failed".into())
}
}
#[test]
fn test_producer_chain_error_isolation() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ProducerInterceptorChain::new(vec![
Arc::new(OrderedProducerInterceptor {
name: "before",
log: Arc::clone(&log),
}),
Arc::new(FailingProducerInterceptor),
Arc::new(OrderedProducerInterceptor {
name: "after",
log: Arc::clone(&log),
}),
]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
chain.close().unwrap();
let log = log.lock().unwrap();
assert_eq!(
*log,
vec![
"before.on_send",
"after.on_send",
"before.close",
"after.close"
]
);
}
#[test]
fn test_consumer_chain_error_isolation() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ConsumerInterceptorChain::new(vec![
Arc::new(OrderedConsumerInterceptor {
name: "before",
log: Arc::clone(&log),
}),
Arc::new(FailingConsumerInterceptor),
Arc::new(OrderedConsumerInterceptor {
name: "after",
log: Arc::clone(&log),
}),
]);
chain.on_consume(&[]).unwrap();
chain.close().unwrap();
let log = log.lock().unwrap();
assert_eq!(
*log,
vec![
"before.on_consume",
"after.on_consume",
"before.close",
"after.close"
]
);
}
#[test]
fn test_safe_wrappers_catch_errors() {
let interceptor = FailingProducerInterceptor;
let mut record = ProducerRecord::new("test", b"v".to_vec());
safe_on_send(&interceptor, &mut record, &mut RecordContext::new());
safe_producer_close(&interceptor);
let interceptor = FailingConsumerInterceptor;
safe_on_consume(&interceptor, &[]);
safe_consumer_close(&interceptor);
}
#[test]
fn test_producer_chain_error_at_first_position() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ProducerInterceptorChain::new(vec![
Arc::new(FailingProducerInterceptor),
Arc::new(OrderedProducerInterceptor {
name: "second",
log: Arc::clone(&log),
}),
]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
let log = log.lock().unwrap();
assert_eq!(*log, vec!["second.on_send"]);
}
#[test]
fn test_producer_chain_error_at_last_position() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ProducerInterceptorChain::new(vec![
Arc::new(OrderedProducerInterceptor {
name: "first",
log: Arc::clone(&log),
}),
Arc::new(FailingProducerInterceptor),
]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
let log = log.lock().unwrap();
assert_eq!(*log, vec!["first.on_send"]);
}
#[test]
fn test_producer_chain_mixed_error_and_panic() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ProducerInterceptorChain::new(vec![
Arc::new(OrderedProducerInterceptor {
name: "first",
log: Arc::clone(&log),
}),
Arc::new(FailingProducerInterceptor),
Arc::new(PanickingProducerInterceptor),
Arc::new(OrderedProducerInterceptor {
name: "last",
log: Arc::clone(&log),
}),
]);
let mut record = ProducerRecord::new("test", b"value".to_vec());
chain
.on_send(&mut record, &mut RecordContext::new())
.unwrap();
let log = log.lock().unwrap();
assert_eq!(*log, vec!["first.on_send", "last.on_send"]);
}
#[test]
fn test_consumer_chain_mixed_error_and_panic() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let chain = ConsumerInterceptorChain::new(vec![
Arc::new(OrderedConsumerInterceptor {
name: "first",
log: Arc::clone(&log),
}),
Arc::new(FailingConsumerInterceptor),
Arc::new(PanickingConsumerInterceptor),
Arc::new(OrderedConsumerInterceptor {
name: "last",
log: Arc::clone(&log),
}),
]);
chain.on_consume(&[]).unwrap();
let log = log.lock().unwrap();
assert_eq!(*log, vec!["first.on_consume", "last.on_consume"]);
}
#[derive(Debug, PartialEq)]
struct Span(&'static str);
#[derive(Debug, PartialEq)]
struct Started(u64);
#[test]
fn record_context_round_trips_distinct_types() {
let mut ctx = RecordContext::new();
assert!(ctx.insert(Span("send")).is_none());
assert!(ctx.insert(Started(7)).is_none());
assert_eq!(ctx.get::<Span>(), Some(&Span("send")));
assert_eq!(ctx.get::<Started>(), Some(&Started(7)));
assert!(ctx.contains::<Span>());
if let Some(started) = ctx.get_mut::<Started>() {
started.0 = 9;
}
assert_eq!(ctx.take::<Started>(), Some(Started(9)));
assert!(
!ctx.contains::<Started>(),
"take removes the value, so a second take sees nothing"
);
assert_eq!(ctx.take::<Started>(), None);
assert_eq!(ctx.get::<Span>(), Some(&Span("send")));
}
#[test]
fn record_context_insert_of_the_same_type_returns_the_previous_value() {
let mut ctx = RecordContext::new();
assert!(ctx.insert(Span("first")).is_none());
assert_eq!(ctx.insert(Span("second")), Some(Span("first")));
assert_eq!(ctx.get::<Span>(), Some(&Span("second")));
}
#[test]
fn record_context_of_a_type_never_stored_is_empty() {
let ctx = RecordContext::new();
assert_eq!(ctx.get::<Span>(), None);
assert!(!ctx.contains::<Span>());
}
#[derive(Debug)]
struct ContextInterceptor {
name: &'static str,
seen: Arc<std::sync::Mutex<Option<Span>>>,
panic_after_store: bool,
}
impl ContextInterceptor {
fn new(name: &'static str) -> (Arc<Self>, Arc<std::sync::Mutex<Option<Span>>>) {
let seen = Arc::new(std::sync::Mutex::new(None));
(
Arc::new(Self {
name,
seen: Arc::clone(&seen),
panic_after_store: false,
}),
seen,
)
}
}
impl ProducerInterceptor for ContextInterceptor {
fn on_send(
&self,
_record: &mut ProducerRecord,
ctx: &mut RecordContext,
) -> InterceptorResult {
ctx.insert(Span(self.name));
if self.panic_after_store {
panic!("interceptor blew up after storing its state");
}
Ok(())
}
fn on_acknowledgement(
&self,
_metadata: &RecordMetadata,
_error: Option<&KrafkaError>,
_headers: &RecordHeaders,
ctx: &mut RecordContext,
) -> InterceptorResult {
*self.seen.lock().unwrap() = ctx.take::<Span>();
Ok(())
}
}
#[test]
fn chained_interceptors_cannot_see_each_others_context() {
let (first, first_seen) = ContextInterceptor::new("first");
let (second, second_seen) = ContextInterceptor::new("second");
let chain = ProducerInterceptorChain::new(vec![first, second]);
let mut ctx = RecordContext::new();
let mut record = ProducerRecord::new("test", b"v".to_vec());
chain.on_send(&mut record, &mut ctx).unwrap();
let metadata = RecordMetadata::failed("test".to_string(), 0);
chain
.on_acknowledgement(&metadata, None, &[], &mut ctx)
.unwrap();
assert_eq!(*first_seen.lock().unwrap(), Some(Span("first")));
assert_eq!(*second_seen.lock().unwrap(), Some(Span("second")));
}
#[test]
fn a_chained_interceptor_cannot_take_a_neighbours_state() {
#[derive(Debug)]
struct Thief {
stole: Arc<std::sync::Mutex<bool>>,
}
impl ProducerInterceptor for Thief {
fn on_send(
&self,
_record: &mut ProducerRecord,
ctx: &mut RecordContext,
) -> InterceptorResult {
*self.stole.lock().unwrap() = ctx.take::<Span>().is_some();
Ok(())
}
}
let (victim, victim_seen) = ContextInterceptor::new("victim");
let stole = Arc::new(std::sync::Mutex::new(false));
let thief = Arc::new(Thief {
stole: Arc::clone(&stole),
});
let chain = ProducerInterceptorChain::new(vec![victim, thief]);
let mut ctx = RecordContext::new();
let mut record = ProducerRecord::new("test", b"v".to_vec());
chain.on_send(&mut record, &mut ctx).unwrap();
let metadata = RecordMetadata::failed("test".to_string(), 0);
chain
.on_acknowledgement(&metadata, None, &[], &mut ctx)
.unwrap();
assert!(!*stole.lock().unwrap(), "the thief must see an empty slot");
assert_eq!(
*victim_seen.lock().unwrap(),
Some(Span("victim")),
"the victim's state must survive the attempt"
);
}
#[test]
fn a_panicking_interceptor_does_not_disturb_its_neighbours_context() {
let (healthy, healthy_seen) = ContextInterceptor::new("healthy");
let exploder = Arc::new(ContextInterceptor {
name: "exploder",
seen: Arc::new(std::sync::Mutex::new(None)),
panic_after_store: true,
});
let chain = ProducerInterceptorChain::new(vec![exploder, healthy]);
let mut ctx = RecordContext::new();
let mut record = ProducerRecord::new("test", b"v".to_vec());
chain.on_send(&mut record, &mut ctx).unwrap();
let metadata = RecordMetadata::failed("test".to_string(), 0);
chain
.on_acknowledgement(&metadata, None, &[], &mut ctx)
.unwrap();
assert_eq!(
*healthy_seen.lock().unwrap(),
Some(Span("healthy")),
"a panic in interceptor 0 must not cost interceptor 1 its state"
);
}
#[test]
#[cfg(target_pointer_width = "64")]
fn a_context_stays_within_its_per_record_size_budget() {
assert_eq!(std::mem::size_of::<RecordContext>(), 32);
}
#[test]
fn an_untouched_context_never_allocates() {
let ctx = RecordContext::new();
assert_eq!(ctx.entries.capacity(), 0, "an empty Vec owns no allocation");
}
#[test]
fn chain_owner_saturates_rather_than_wrapping() {
assert_eq!(chain_owner(0), 0);
assert_eq!(chain_owner(usize::from(u16::MAX)), u16::MAX);
assert_eq!(
chain_owner(usize::from(u16::MAX) + 1),
u16::MAX,
"an absurd chain collapses its tail rather than aliasing onto index 0"
);
}
}