use std::collections::HashMap;
use std::sync::{Arc, Mutex, Weak};
use std::time::Duration;
use futures::future::select_all;
use rdkafka::consumer::{Consumer as _, ConsumerGroupMetadata, StreamConsumer};
use rdkafka::{Offset, TopicPartitionList};
use ruststream::codec::Codec;
#[cfg(any(feature = "json", feature = "cbor", feature = "msgpack"))]
use ruststream::codec::DefaultCodec;
use ruststream::runtime::{
Outgoing, PublishContext, PublishTransform, PublishTransformIdentity, PublishTransformStack,
TypedPublisher,
};
use ruststream::{OutgoingMessage, Publisher, TransactionalPublisher as _};
use tracing::{debug, error};
use crate::error::KafkaError;
use crate::publisher::KafkaPublisher;
use crate::tracker::{CommitTracker, TrackingContext};
const DEFAULT_COMMIT_INTERVAL: Duration = Duration::from_millis(100);
#[derive(Clone)]
pub(crate) struct EosSource {
tracker: Weak<CommitTracker>,
consumer: Weak<StreamConsumer<TrackingContext>>,
}
impl EosSource {
pub(crate) fn new(
tracker: &Arc<CommitTracker>,
consumer: &Arc<StreamConsumer<TrackingContext>>,
) -> Self {
Self {
tracker: Arc::downgrade(tracker),
consumer: Arc::downgrade(consumer),
}
}
pub(crate) fn alive(&self) -> bool {
self.tracker.strong_count() > 0 && self.consumer.strong_count() > 0
}
fn upgrade(&self) -> Option<LiveSource> {
Some(LiveSource {
tracker: self.tracker.upgrade()?,
consumer: self.consumer.upgrade()?,
})
}
}
struct LiveSource {
tracker: Arc<CommitTracker>,
consumer: Arc<StreamConsumer<TrackingContext>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SourceOffset {
topic: String,
partition: i32,
offset: i64,
}
impl SourceOffset {
#[must_use]
pub fn new(topic: impl Into<String>, partition: i32, offset: i64) -> Self {
Self {
topic: topic.into(),
partition,
offset,
}
}
fn key(&self) -> (String, i32) {
(self.topic.clone(), self.partition)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Phase {
Idle,
Opening,
Open,
Committing,
}
#[derive(Debug)]
struct Window {
phase: Phase,
enrolled: HashMap<(String, i32), i64>,
failed: bool,
epoch: u64,
}
struct PipelineInner {
publisher: KafkaPublisher,
id: Option<String>,
interval: Duration,
window: Mutex<Window>,
phase_changed: tokio::sync::Notify,
committed: Mutex<HashMap<(String, i32), i64>>,
session_low: Mutex<HashMap<(String, i32), i64>>,
}
#[derive(Clone)]
pub struct EosPipeline {
inner: Arc<PipelineInner>,
}
impl std::fmt::Debug for EosPipeline {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EosPipeline")
.field("id", &self.inner.id)
.field("interval", &self.inner.interval)
.finish_non_exhaustive()
}
}
impl EosPipeline {
#[must_use]
pub fn new(publisher: KafkaPublisher) -> Self {
let id = publisher.transactional_id_str().map(str::to_owned);
Self {
inner: Arc::new(PipelineInner {
publisher,
id,
interval: DEFAULT_COMMIT_INTERVAL,
window: Mutex::new(Window {
phase: Phase::Idle,
enrolled: HashMap::new(),
failed: false,
epoch: 0,
}),
phase_changed: tokio::sync::Notify::new(),
committed: Mutex::new(HashMap::new()),
session_low: Mutex::new(HashMap::new()),
}),
}
}
#[must_use]
pub fn commit_interval(self, interval: Duration) -> Self {
Self {
inner: Arc::new(PipelineInner {
publisher: self.inner.publisher.clone(),
id: self.inner.id.clone(),
interval,
window: Mutex::new(Window {
phase: Phase::Idle,
enrolled: HashMap::new(),
failed: false,
epoch: 0,
}),
phase_changed: tokio::sync::Notify::new(),
committed: Mutex::new(HashMap::new()),
session_low: Mutex::new(HashMap::new()),
}),
}
}
pub async fn publish(
&self,
source: &SourceOffset,
msg: OutgoingMessage<'_>,
) -> Result<(), KafkaError> {
self.admit(source).await?;
let sent = self.inner.publisher.publish(msg).await;
if sent.is_err() {
let mut window = self.inner.window.lock().expect("window mutex poisoned");
window.failed = true;
}
sent
}
async fn admit(&self, source: &SourceOffset) -> Result<(), KafkaError> {
loop {
let phase_changed = self.inner.phase_changed.notified();
let action = {
let mut window = self.inner.window.lock().expect("window mutex poisoned");
match window.phase {
Phase::Open => {
Self::enroll(&mut window, &self.inner.session_low, source);
Admission::Admitted
}
Phase::Idle => {
window.phase = Phase::Opening;
Admission::Opener
}
Phase::Opening => Admission::Wait,
Phase::Committing => {
let participant = window
.enrolled
.get(&source.key())
.is_some_and(|max| source.offset <= *max);
if participant {
Admission::Admitted
} else {
Admission::Wait
}
}
}
};
match action {
Admission::Admitted => return Ok(()),
Admission::Wait => {
phase_changed.await;
}
Admission::Opener => return self.open_window(source).await,
}
}
}
async fn open_window(&self, source: &SourceOffset) -> Result<(), KafkaError> {
let begun = self.inner.publisher.begin_transaction().await;
let mut window = self.inner.window.lock().expect("window mutex poisoned");
match begun {
Ok(()) => {
window.phase = Phase::Open;
window.failed = false;
Self::enroll(&mut window, &self.inner.session_low, source);
let epoch = window.epoch;
drop(window);
tokio::spawn(run_window(Arc::clone(&self.inner), epoch));
}
Err(err) => {
window.phase = Phase::Idle;
drop(window);
self.inner.phase_changed.notify_waiters();
return Err(err);
}
}
self.inner.phase_changed.notify_waiters();
Ok(())
}
fn enroll(
window: &mut Window,
session_low: &Mutex<HashMap<(String, i32), i64>>,
source: &SourceOffset,
) {
let key = source.key();
session_low
.lock()
.expect("session low mutex poisoned")
.entry(key.clone())
.or_insert(source.offset);
let max = window.enrolled.entry(key).or_insert(source.offset);
if source.offset > *max {
*max = source.offset;
}
}
}
enum Admission {
Admitted,
Wait,
Opener,
}
async fn run_window(inner: Arc<PipelineInner>, epoch: u64) {
tokio::time::sleep(inner.interval).await;
let enrolled = {
let mut window = inner.window.lock().expect("window mutex poisoned");
if window.epoch != epoch || window.phase != Phase::Open {
return;
}
window.phase = Phase::Committing;
window.enrolled.clone()
};
let outcome = commit_window(&inner, &enrolled).await;
{
let mut window = inner.window.lock().expect("window mutex poisoned");
window.phase = Phase::Idle;
window.enrolled.clear();
window.failed = false;
window.epoch += 1;
}
inner.phase_changed.notify_waiters();
if let Err(err) = outcome {
error!(
target: "ruststream_rdkafka",
pipeline = inner.id.as_deref().unwrap_or("<no id>"),
error = %err,
"EOS window aborted; its sources seek back and the window redelivers",
);
}
}
async fn commit_window(
inner: &Arc<PipelineInner>,
enrolled: &HashMap<(String, i32), i64>,
) -> Result<(), KafkaError> {
let id = inner.id.clone().ok_or_else(|| {
KafkaError::InvalidOptions(
"an EosPipeline publisher needs `KafkaPublisher::transactional_id`".to_owned(),
)
})?;
let conn = inner.publisher.shared_conn();
let state = conn.get().ok_or(KafkaError::NotConnected)?;
let sources: Vec<LiveSource> = state
.eos_sources(&id)
.iter()
.filter_map(EosSource::upgrade)
.collect();
let failed = {
let window = inner.window.lock().expect("window mutex poisoned");
window.failed
};
let ready = if failed {
Err(KafkaError::Publish(
"a publish into this window failed; the transaction is poisoned"
.to_owned()
.into(),
))
} else {
wait_settled(inner, &sources, enrolled).await
};
let result = match ready {
Ok(()) => try_commit(inner, &sources).await,
Err(err) => Err(err),
};
if let Err(err) = result {
abort_window(inner, &sources, enrolled).await;
return Err(err);
}
Ok(())
}
async fn wait_settled(
inner: &Arc<PipelineInner>,
sources: &[LiveSource],
enrolled: &HashMap<(String, i32), i64>,
) -> Result<(), KafkaError> {
let deadline = tokio::time::Instant::now() + inner.publisher.transaction_deadline();
loop {
let waiters: Vec<_> = sources
.iter()
.map(|source| Box::pin(source.tracker.advance_waiter()))
.collect();
let pending = enrolled.iter().find(|((topic, partition), max)| {
!sources.iter().any(|source| {
source
.tracker
.stored_position(topic, *partition)
.is_some_and(|stored| stored >= **max)
})
});
let Some(((topic, partition), max)) = pending else {
return Ok(());
};
if waiters.is_empty() {
return Err(KafkaError::InvalidOptions(format!(
"EOS pipeline has no registered sources for its id; is the subscription in \
`Commit::Transactional` mode with the matching pipeline id? (waiting on \
{topic}[{partition}] up to offset {max})",
)));
}
debug!(
target: "ruststream_rdkafka",
topic = %topic,
partition = partition,
up_to = max,
"EOS window waiting for participants to settle",
);
if tokio::time::timeout_at(deadline, select_all(waiters))
.await
.is_err()
{
return Err(KafkaError::Publish(
format!(
"EOS window stalled: {topic}[{partition}] did not settle up to offset \
{max} within the transaction deadline (a hung or retrying handler, or a \
revoked partition)",
)
.into(),
));
}
}
}
async fn try_commit(inner: &Arc<PipelineInner>, sources: &[LiveSource]) -> Result<(), KafkaError> {
let mut sent: Vec<((String, i32), i64)> = Vec::new();
for source in sources {
let positions = source.tracker.stored_positions();
if positions.is_empty() {
continue;
}
let mut offsets = TopicPartitionList::new();
for ((topic, partition), stored) in &positions {
offsets
.add_partition_offset(topic, *partition, Offset::Offset(stored + 1))
.map_err(KafkaError::publish)?;
}
let metadata = group_metadata(source)?;
inner.publisher.send_offsets(offsets, metadata).await?;
sent.extend(positions.into_iter().map(|(key, stored)| (key, stored + 1)));
}
inner.publisher.commit().await?;
{
let mut committed = inner.committed.lock().expect("committed mutex poisoned");
for (key, next) in sent {
committed.insert(key, next);
}
}
Ok(())
}
fn group_metadata(source: &LiveSource) -> Result<ConsumerGroupMetadata, KafkaError> {
source.consumer.group_metadata().ok_or_else(|| {
KafkaError::Publish(
"the source consumer has no group metadata (not a group member yet or already \
closed); cannot commit its offsets transactionally"
.to_owned()
.into(),
)
})
}
async fn abort_window(
inner: &Arc<PipelineInner>,
sources: &[LiveSource],
enrolled: &HashMap<(String, i32), i64>,
) {
if let Err(err) = inner.publisher.abort().await {
error!(
target: "ruststream_rdkafka",
error = %err,
"EOS window abort failed; the transaction resolves by its broker-side timeout",
);
}
let committed = inner
.committed
.lock()
.expect("committed mutex poisoned")
.clone();
let session_low = inner
.session_low
.lock()
.expect("session low mutex poisoned")
.clone();
for key @ (topic, partition) in enrolled.keys() {
let Some(target) = committed
.get(key)
.copied()
.or_else(|| session_low.get(key).copied())
else {
continue;
};
let Some(source) = sources
.iter()
.find(|source| source.tracker.covers(topic, *partition))
else {
continue;
};
if let Err(err) = source.consumer.seek(
topic,
*partition,
Offset::Offset(target),
Duration::from_secs(5),
) {
debug!(
target: "ruststream_rdkafka",
topic = %topic,
partition = partition,
error = %err,
"seek-back after an aborted EOS window failed",
);
}
}
}
pub const EOS_SOURCE_HEADER: &str = "kafka-eos-source";
pub(crate) fn encode_source(topic: &str, partition: i32, offset: i64) -> String {
format!("{partition}:{offset}:{topic}")
}
fn decode_source(value: &str) -> Option<SourceOffset> {
let mut parts = value.splitn(3, ':');
let partition = parts.next()?.parse().ok()?;
let offset = parts.next()?.parse().ok()?;
let topic = parts.next()?;
Some(SourceOffset::new(topic, partition, offset))
}
#[derive(Debug, Clone, Copy, Default)]
pub struct EosReplies;
impl<C> PublishTransform<C> for EosReplies {
fn apply(&self, out: &mut Outgoing<'_>, cx: &PublishContext<'_, C>) {
if let Some(source) = cx.headers().get(EOS_SOURCE_HEADER) {
let source = source.to_vec();
out.headers_mut().insert(EOS_SOURCE_HEADER, source);
}
}
}
impl EosPipeline {
#[cfg(any(feature = "json", feature = "cbor", feature = "msgpack"))]
#[must_use]
pub fn replies(
&self,
) -> TypedPublisher<
Self,
DefaultCodec,
PublishTransformStack<PublishTransformIdentity, EosReplies>,
> {
TypedPublisher::new(self.clone()).transform(EosReplies)
}
#[must_use]
pub fn replies_with<C: Codec>(
&self,
codec: C,
) -> TypedPublisher<Self, C, PublishTransformStack<PublishTransformIdentity, EosReplies>> {
TypedPublisher::with_codec(self.clone(), codec).transform(EosReplies)
}
}
impl Publisher for EosPipeline {
type Error = KafkaError;
async fn publish(&self, msg: OutgoingMessage<'_>) -> Result<(), Self::Error> {
let Some(source) = msg
.headers()
.get_str(EOS_SOURCE_HEADER)
.and_then(decode_source)
else {
return Err(KafkaError::InvalidOptions(
"an EOS reply carries no source coordinates: the subscription must be in \
`Commit::Transactional` mode for this pipeline, and the reply publisher must \
relay them (wire it with `EosPipeline::replies()` or add the `EosReplies` \
transform)"
.to_owned(),
));
};
let mut headers = msg.headers().clone();
headers.remove(EOS_SOURCE_HEADER);
let stripped = OutgoingMessage::new(msg.name(), msg.payload()).with_headers(headers);
self.publish(&source, stripped).await
}
}
#[cfg(test)]
mod tests {
use ruststream::Headers;
use super::*;
#[test]
fn source_header_roundtrips_topics_with_colons() {
let encoded = encode_source("orders:eu:v1", 3, 42);
let decoded = decode_source(&encoded).expect("decodes");
assert_eq!(decoded, SourceOffset::new("orders:eu:v1", 3, 42));
}
#[test]
fn malformed_source_headers_are_rejected() {
for bad in ["", "3", "3:x:orders", "x:42:orders"] {
assert!(decode_source(bad).is_none(), "{bad:?} must not decode");
}
}
#[tokio::test]
async fn reply_without_source_coordinates_fails_clearly() {
let pipeline = EosPipeline::new(KafkaPublisher::new(Arc::default()).transactional_id("p1"));
let err = Publisher::publish(&pipeline, OutgoingMessage::new("replies", b"x".as_slice()))
.await
.expect_err("a reply without the source header must fail");
assert!(matches!(err, KafkaError::InvalidOptions(_)));
assert!(err.to_string().contains("Commit::Transactional"));
let _ = Headers::new();
}
}