use std::fmt;
use std::sync::Arc;
use futures::Stream;
use ruststream::testing::Coordinator;
use ruststream::{AckError, BatchSubscriber, Headers, IncomingMessage, Partitioned, Subscriber};
use super::broker::TestBrokerState;
use super::router::{DeliveryReceiver, DeliverySender, SubscriptionId, TestDelivery};
use crate::error::KafkaError;
pub struct KafkaTestSubscriber {
state: Arc<TestBrokerState>,
ids: Vec<SubscriptionId>,
topic: String,
sender: DeliverySender,
receiver: DeliveryReceiver,
coordinator: Option<Coordinator>,
}
impl KafkaTestSubscriber {
pub(crate) fn open_many(state: &Arc<TestBrokerState>, topics: &[String]) -> Self {
let (ids, sender, receiver) = state.router.subscribe_many(topics);
let coordinator = state.coordinator();
Self {
state: Arc::clone(state),
ids,
topic: topics.join(","),
sender,
receiver,
coordinator,
}
}
#[must_use]
pub fn topic(&self) -> &str {
&self.topic
}
}
impl Drop for KafkaTestSubscriber {
fn drop(&mut self) {
for id in &self.ids {
self.state.router.unsubscribe(*id);
}
}
}
impl fmt::Debug for KafkaTestSubscriber {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("KafkaTestSubscriber")
.field("topic", &self.topic)
.finish_non_exhaustive()
}
}
impl Subscriber for KafkaTestSubscriber {
type Message = KafkaTestMessage;
type Error = KafkaError;
fn stream(&mut self) -> impl Stream<Item = Result<Self::Message, Self::Error>> + Send + '_ {
let Self {
receiver,
sender,
coordinator,
..
} = self;
futures::stream::poll_fn(move |cx| {
receiver.poll_recv(cx).map(|delivery| {
delivery.map(|delivery| {
Ok(KafkaTestMessage {
delivery: Some(delivery),
sender: sender.clone(),
coordinator: coordinator.clone(),
})
})
})
})
}
}
impl BatchSubscriber for KafkaTestSubscriber {
type Batch = Vec<KafkaTestMessage>;
fn batches(
&mut self,
) -> impl Stream<Item = Result<Self::Batch, <Self as Subscriber>::Error>> + Send + '_ {
let Self {
receiver,
sender,
coordinator,
..
} = self;
futures::stream::poll_fn(move |cx| {
receiver.poll_recv(cx).map(|delivery| {
delivery.map(|first| {
let mut batch = vec![KafkaTestMessage {
delivery: Some(first),
sender: sender.clone(),
coordinator: coordinator.clone(),
}];
while let Ok(delivery) = receiver.try_recv() {
batch.push(KafkaTestMessage {
delivery: Some(delivery),
sender: sender.clone(),
coordinator: coordinator.clone(),
});
}
Ok(batch)
})
})
})
}
}
pub struct KafkaTestMessage {
delivery: Option<TestDelivery>,
sender: DeliverySender,
coordinator: Option<Coordinator>,
}
impl KafkaTestMessage {
fn take(&mut self) -> TestDelivery {
self.delivery
.take()
.expect("KafkaTestMessage settled twice")
}
}
impl Drop for KafkaTestMessage {
fn drop(&mut self) {
if let Some(coordinator) = self.coordinator.take() {
coordinator.consumed();
}
}
}
impl fmt::Debug for KafkaTestMessage {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("KafkaTestMessage")
.field("delivery", &self.delivery)
.finish_non_exhaustive()
}
}
impl IncomingMessage for KafkaTestMessage {
fn payload(&self) -> &[u8] {
&self
.delivery
.as_ref()
.expect("message accessed after settlement")
.payload
}
fn headers(&self) -> &Headers {
&self
.delivery
.as_ref()
.expect("message accessed after settlement")
.headers
}
fn partition_key(&self) -> Option<&[u8]> {
self.headers().get(crate::PARTITION_KEY_HEADER)
}
async fn ack(mut self) -> Result<(), AckError> {
drop(self.take());
Ok(())
}
async fn nack(mut self, requeue: bool) -> Result<(), AckError> {
let delivery = self.take();
if requeue && self.sender.send(delivery).is_ok() {
if let Some(coordinator) = &self.coordinator {
coordinator.enqueued();
}
}
Ok(())
}
}
impl Partitioned for KafkaTestMessage {
fn partition_key(&self) -> Option<&[u8]> {
self.headers().get(crate::PARTITION_KEY_HEADER)
}
}