use std::collections::HashMap;
use std::fmt;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use rdkafka::consumer::{Consumer as _, StreamConsumer};
use rdkafka::producer::{FutureProducer, Producer as _};
use rdkafka::{ClientConfig, Offset, TopicPartitionList};
use ruststream::{Broker, DescribeServer, ServerSpec, Subscribe};
use tokio::sync::OnceCell;
use tokio::task;
use crate::eos::EosSource;
use crate::error::KafkaError;
use crate::publisher::KafkaPublisher;
use crate::retry::RetryContext;
use crate::subscriber::KafkaSubscriber;
use crate::topic::{Commit, KafkaTopic, StartOffset};
use crate::tracker::{CommitTracker, TrackingContext};
pub(crate) struct ConnState {
producer: FutureProducer,
producer_config: ClientConfig,
eos_sources: Mutex<HashMap<String, Vec<EosSource>>>,
}
impl ConnState {
pub(crate) fn producer(&self) -> &FutureProducer {
&self.producer
}
pub(crate) fn producer_config(&self) -> &ClientConfig {
&self.producer_config
}
fn register_eos(&self, pipeline: &str, source: EosSource) {
let mut sources = self
.eos_sources
.lock()
.expect("eos source registry mutex poisoned");
sources.entry(pipeline.to_owned()).or_default().push(source);
}
pub(crate) fn eos_sources(&self, pipeline: &str) -> Vec<EosSource> {
let mut sources = self
.eos_sources
.lock()
.expect("eos source registry mutex poisoned");
sources
.get_mut(pipeline)
.map_or_else(Vec::new, |registered| {
registered.retain(EosSource::alive);
registered.clone()
})
}
}
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,
#[cfg(feature = "schema-registry")]
schema_registry: Option<crate::schema_registry::SchemaRegistry>,
}
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,
#[cfg(feature = "schema-registry")]
schema_registry: None,
}
}
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
}
#[cfg(feature = "schema-registry")]
#[must_use]
pub fn schema_registry(mut self, registry: crate::schema_registry::SchemaRegistry) -> Self {
self.schema_registry = Some(registry);
self
}
#[allow(clippy::unused_async)]
pub async fn subscribe(&self, def: KafkaTopic) -> Result<KafkaSubscriber, KafkaError> {
self.connected()?;
def.validate()?;
let manual = !def.assigned_partitions().is_empty();
let group = def
.group_or(self.default_group.as_deref())
.map(str::to_owned);
let group = match group {
Some(group) => Some(group),
None if manual => None,
None => {
return Err(KafkaError::InvalidOptions(format!(
"subscription to {:?} has no consumer group: set `KafkaTopic::group` or \
`KafkaBroker::default_group`",
def.topic(),
)));
}
};
if manual {
validate_manual_assignment(&def, group.as_deref())?;
}
let mut config = self.base_config();
if let Some(group) = &group {
config.set("group.id", group);
} else {
config.set("group.id", "ruststream.standalone");
config.set("enable.auto.commit", "false");
}
match def.start_offset() {
StartOffset::Committed => {}
StartOffset::Earliest => {
config.set("auto.offset.reset", "earliest");
}
StartOffset::Latest => {
config.set("auto.offset.reset", "latest");
}
}
if let Some(assignment) = def.assignment_strategy() {
config.set(
"partition.assignment.strategy",
assignment.as_config_value(),
);
}
match def.commit_mode() {
Commit::Auto => {}
Commit::Tracked => {
config.set("enable.auto.offset.store", "false");
}
Commit::Transactional(_) => {
config.set("enable.auto.offset.store", "false");
config.set("enable.auto.commit", "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)?;
if manual {
assign_partitions(&consumer, &def)?;
} else {
let names: Vec<&str> = def.subscribed_topics().iter().map(String::as_str).collect();
consumer.subscribe(&names).map_err(KafkaError::subscribe)?;
}
let consumer = Arc::new(consumer);
if let Commit::Transactional(pipeline) = def.commit_mode() {
let state = self.conn.get().ok_or(KafkaError::NotConnected)?;
state.register_eos(pipeline, EosSource::new(&tracker, &consumer));
}
let retry =
(def.retry_policy().is_some() || def.dead_letter_topic().is_some()).then(|| {
Arc::new(RetryContext::new(
def.retry_policy().cloned(),
def.max_deliveries_cap(),
def.dead_letter_topic().map(str::to_owned),
Arc::clone(&self.conn),
Arc::clone(&consumer),
))
});
let subscriber = KafkaSubscriber::new(
consumer,
def.topic().to_owned(),
def.commit_mode().clone(),
tracker,
def.lane_key_choice(),
retry,
);
#[cfg(feature = "schema-registry")]
let subscriber = subscriber.with_schema_registry(self.schema_registry.clone());
Ok(subscriber)
}
}
fn validate_manual_assignment(def: &KafkaTopic, group: Option<&str>) -> Result<(), KafkaError> {
if matches!(def.commit_mode(), Commit::Transactional(_)) {
return Err(KafkaError::InvalidOptions(
"manual partition assignment does not compose with `Commit::Transactional`: an \
EOS pipeline commits through the consumer group protocol"
.to_owned(),
));
}
if group.is_none() {
if def.commit_mode() == &Commit::Tracked {
return Err(KafkaError::InvalidOptions(
"`Commit::Tracked` needs a group to commit into; name one with \
`KafkaTopic::group` or drop the commit mode for a group-less reader"
.to_owned(),
));
}
if def.start_offset() == StartOffset::Committed {
return Err(KafkaError::InvalidOptions(
"a group-less manual assignment has no committed offsets to start from; set \
`start(StartOffset::Earliest)` or `Latest`, or name a group"
.to_owned(),
));
}
}
Ok(())
}
fn assign_partitions(
consumer: &StreamConsumer<TrackingContext>,
def: &KafkaTopic,
) -> Result<(), KafkaError> {
let offset = match def.start_offset() {
StartOffset::Committed => Offset::Stored,
StartOffset::Earliest => Offset::Beginning,
StartOffset::Latest => Offset::End,
};
let mut assignment = TopicPartitionList::new();
for partition in def.assigned_partitions() {
assignment
.add_partition_offset(def.topic(), *partition, offset)
.map_err(KafkaError::subscribe)?;
}
consumer.assign(&assignment).map_err(KafkaError::subscribe)
}
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,
producer_config: config,
eos_sources: Mutex::new(HashMap::new()),
})
})
.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(_)));
}
}