use std::{collections::HashMap, fmt, sync::Arc};
use super::ConsumerRecords;
use crate::common::{OffsetAndMetadata, TopicPartition};
#[derive(Debug, Clone, Default)]
pub struct InterceptorConfigs {
pub client_id: Option<String>,
pub group_id: Option<String>,
}
pub trait ConsumerInterceptor: Send + Sync + 'static {
fn configure(&self, _configs: &InterceptorConfigs) {}
fn on_consume(&self, records: ConsumerRecords) -> ConsumerRecords {
records
}
fn on_commit(&self, _offsets: &HashMap<TopicPartition, OffsetAndMetadata>) {}
fn close(&self) {}
}
#[derive(Clone, Default)]
pub(super) struct ConsumerInterceptors {
inner: Arc<[Arc<dyn ConsumerInterceptor>]>,
}
impl ConsumerInterceptors {
pub(super) fn push_and_configure(
&mut self,
interceptor: impl ConsumerInterceptor,
configs: &InterceptorConfigs,
) {
let interceptor: Arc<dyn ConsumerInterceptor> = Arc::new(interceptor);
let configured = Arc::clone(&interceptor);
let _ignored = catch_interceptor_unwind(|| configured.configure(configs));
let mut inner = self.inner.to_vec();
inner.push(interceptor);
self.inner = Arc::from(inner.into_boxed_slice());
}
pub(super) fn on_consume(&self, mut records: ConsumerRecords) -> ConsumerRecords {
for interceptor in self.inner.iter() {
let previous = records.clone();
match catch_interceptor_unwind(|| interceptor.on_consume(records)) {
Some(intercepted) => records = intercepted,
None => records = previous,
}
}
records
}
pub(super) fn on_commit(&self, offsets: &HashMap<TopicPartition, OffsetAndMetadata>) {
for interceptor in self.inner.iter() {
let _ignored = catch_interceptor_unwind(|| interceptor.on_commit(offsets));
}
}
pub(super) fn close(&self) {
for interceptor in self.inner.iter() {
let _ignored = catch_interceptor_unwind(|| interceptor.close());
}
}
}
fn catch_interceptor_unwind<T>(f: impl FnOnce() -> T) -> Option<T> {
std::panic::catch_unwind(std::panic::AssertUnwindSafe(f)).ok()
}
impl fmt::Debug for ConsumerInterceptors {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConsumerInterceptors")
.field("len", &self.inner.len())
.finish()
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
#[derive(Default)]
struct Recorder {
commits: Mutex<usize>,
configured: Mutex<Vec<Option<String>>>,
}
impl ConsumerInterceptor for Arc<Recorder> {
fn configure(&self, configs: &InterceptorConfigs) {
self.configured
.lock()
.unwrap()
.push(configs.group_id.clone());
}
fn on_consume(&self, records: ConsumerRecords) -> ConsumerRecords {
let _dropped = records;
ConsumerRecords::empty()
}
fn on_commit(&self, _offsets: &HashMap<TopicPartition, OffsetAndMetadata>) {
let mut commits = self.commits.lock().unwrap();
*commits = commits.saturating_add(1);
}
}
#[test]
fn chain_threads_records_and_fires_commit() {
let recorder = Arc::new(Recorder::default());
let mut chain = ConsumerInterceptors::default();
chain.push_and_configure(
Arc::clone(&recorder),
&InterceptorConfigs {
client_id: None,
group_id: Some("g".to_owned()),
},
);
assert_eq!(
*recorder.configured.lock().unwrap(),
vec![Some("g".to_owned())]
);
let mut records = ConsumerRecords::empty();
records.push_partition("t".to_owned(), 0, Vec::new());
assert!(chain.on_consume(records).is_empty());
chain.on_commit(&HashMap::new());
assert_eq!(*recorder.commits.lock().unwrap(), 1);
}
}