use std::fmt;
use std::sync::Arc;
use std::time::Duration;
use rdkafka::ClientConfig;
use rdkafka::consumer::{Consumer as _, StreamConsumer};
use rdkafka::producer::{FutureProducer, Producer as _};
use ruststream::{Broker, DescribeServer, ServerSpec, Subscribe};
use tokio::sync::OnceCell;
use tokio::task;
use crate::error::KafkaError;
use crate::publisher::KafkaPublisher;
use crate::subscriber::KafkaSubscriber;
use crate::topic::{Commit, KafkaTopic, StartOffset};
use crate::tracker::{CommitTracker, TrackingContext};
pub(crate) struct ConnState {
producer: FutureProducer,
}
impl ConnState {
pub(crate) fn producer(&self) -> &FutureProducer {
&self.producer
}
}
impl fmt::Debug for ConnState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConnState").finish_non_exhaustive()
}
}
pub(crate) type SharedConn = Arc<OnceCell<ConnState>>;
const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
const DEFAULT_FLUSH_TIMEOUT: Duration = Duration::from_secs(30);
#[derive(Debug, Clone)]
pub struct KafkaBroker {
conn: SharedConn,
servers: Vec<String>,
default_group: Option<String>,
client_config: Vec<(String, String)>,
producer_config: Vec<(String, String)>,
connect_timeout: Duration,
flush_timeout: Duration,
}
impl KafkaBroker {
#[must_use]
pub fn new<I, S>(servers: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
Self {
conn: Arc::new(OnceCell::new()),
servers: servers.into_iter().map(Into::into).collect(),
default_group: None,
client_config: Vec::new(),
producer_config: Vec::new(),
connect_timeout: DEFAULT_CONNECT_TIMEOUT,
flush_timeout: DEFAULT_FLUSH_TIMEOUT,
}
}
pub async fn connect<I, S>(servers: I) -> Result<Self, KafkaError>
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
let broker = Self::new(servers);
Broker::connect(&broker).await?;
Ok(broker)
}
#[must_use]
pub fn default_group(mut self, group: impl Into<String>) -> Self {
self.default_group = Some(group.into());
self
}
#[must_use]
pub fn config(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.client_config.push((key.into(), value.into()));
self
}
#[must_use]
pub fn producer_config(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.producer_config.push((key.into(), value.into()));
self
}
#[must_use]
pub fn connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout;
self
}
#[must_use]
pub fn flush_timeout(mut self, timeout: Duration) -> Self {
self.flush_timeout = timeout;
self
}
#[must_use]
pub fn publisher(&self) -> KafkaPublisher {
KafkaPublisher::new(Arc::clone(&self.conn))
}
fn connected(&self) -> Result<&ConnState, KafkaError> {
self.conn.get().ok_or(KafkaError::NotConnected)
}
fn base_config(&self) -> ClientConfig {
let mut config = ClientConfig::new();
config.set("bootstrap.servers", self.servers.join(","));
for (key, value) in &self.client_config {
config.set(key, value);
}
config
}
#[allow(clippy::unused_async)]
pub async fn subscribe(&self, def: KafkaTopic) -> Result<KafkaSubscriber, KafkaError> {
self.connected()?;
let group = def
.group_or(self.default_group.as_deref())
.ok_or_else(|| {
KafkaError::InvalidOptions(format!(
"subscription to {:?} has no consumer group: set `KafkaTopic::group` or \
`KafkaBroker::default_group`",
def.topic(),
))
})?
.to_owned();
let mut config = self.base_config();
config.set("group.id", group);
match def.start_offset() {
StartOffset::Committed => {}
StartOffset::Earliest => {
config.set("auto.offset.reset", "earliest");
}
StartOffset::Latest => {
config.set("auto.offset.reset", "latest");
}
}
if def.commit_mode() == Commit::Tracked {
config.set("enable.auto.offset.store", "false");
}
for (key, value) in def.config_entries() {
config.set(key, value);
}
let tracker = Arc::new(CommitTracker::default());
let context = TrackingContext::new(Arc::clone(&tracker));
let consumer: StreamConsumer<TrackingContext> = config
.create_with_context(context)
.map_err(KafkaError::subscribe)?;
consumer
.subscribe(&[def.topic()])
.map_err(KafkaError::subscribe)?;
Ok(KafkaSubscriber::new(
Arc::new(consumer),
def.topic().to_owned(),
def.commit_mode(),
tracker,
))
}
}
impl Broker for KafkaBroker {
type Error = KafkaError;
async fn connect(&self) -> Result<(), Self::Error> {
self.conn
.get_or_try_init(|| async {
if self.servers.is_empty() {
return Err(KafkaError::InvalidOptions(
"at least one bootstrap server is required".to_owned(),
));
}
let mut config = self.base_config();
for (key, value) in &self.producer_config {
config.set(key, value);
}
let producer: FutureProducer = config.create().map_err(KafkaError::connect)?;
let probe = producer.clone();
let timeout = self.connect_timeout;
task::spawn_blocking(move || probe.client().fetch_metadata(None, timeout))
.await
.map_err(|err| KafkaError::Connect(Box::new(err)))?
.map_err(KafkaError::connect)?;
Ok(ConnState { producer })
})
.await?;
Ok(())
}
async fn shutdown(&self) -> Result<(), Self::Error> {
if let Some(state) = self.conn.get() {
let producer = state.producer.clone();
let timeout = self.flush_timeout;
task::spawn_blocking(move || producer.flush(timeout))
.await
.map_err(|err| KafkaError::Publish(Box::new(err)))?
.map_err(KafkaError::publish)?;
}
Ok(())
}
}
#[allow(clippy::use_self)]
impl Subscribe for KafkaBroker {
type Subscriber = KafkaSubscriber;
async fn subscribe(&self, name: &str) -> Result<Self::Subscriber, Self::Error> {
KafkaBroker::subscribe(self, KafkaTopic::new(name)).await
}
}
impl DescribeServer for KafkaBroker {
fn describe_server(&self) -> ServerSpec {
ServerSpec::new(self.servers.join(","), "kafka")
}
}
#[cfg(test)]
mod tests {
use ruststream::DescribeServer as _;
use super::*;
#[test]
fn construction_is_synchronous_and_io_free() {
let broker = KafkaBroker::new(["a:9092", "b:9092"]).default_group("g");
assert_eq!(
broker.describe_server().host.as_deref(),
Some("a:9092,b:9092")
);
assert_eq!(broker.describe_server().protocol, "kafka");
}
#[tokio::test]
async fn operations_before_connect_report_not_connected() {
let broker = KafkaBroker::new(["localhost:9092"]).default_group("g");
let err = broker
.subscribe(KafkaTopic::new("orders"))
.await
.unwrap_err();
assert!(matches!(err, KafkaError::NotConnected));
}
#[tokio::test]
async fn missing_group_is_a_clear_startup_error() {
let broker = KafkaBroker::new(["localhost:9092"]);
let def = KafkaTopic::new("orders");
assert!(def.group_or(None).is_none());
assert_eq!(def.group_or(Some("fallback")), Some("fallback"));
let _ = broker;
}
#[tokio::test]
async fn connect_with_no_servers_fails_fast() {
let broker = KafkaBroker::new(Vec::<String>::new());
let err = Broker::connect(&broker).await.unwrap_err();
assert!(matches!(err, KafkaError::InvalidOptions(_)));
}
}