use crate::sink::context::SinkContext;
use crate::sink::encoder::{
DEFAULT_MAX_MESSAGE_BYTES, KafkaBytesEncoder, KafkaEncoder, KafkaJsonEncoder, MessageEncoder,
};
use crate::sink::writer::{KafkaEndpoint, KafkaWriter};
use bytesize::ByteSize;
use rdkafka::producer::ThreadedProducer;
use serde::Deserialize;
use spate_core::config::{ComponentConfig, ConfigError};
use spate_core::deser::{Owned, RecFamily};
use spate_core::sink::{
BatchConfig, BreakerConfig, InflightConfig, RetryConfig, SinkBundle, SinkParts, SinkPoolConfig,
SinkProbeFn, endpoint_probe,
};
use std::collections::BTreeMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
const DENYLIST: &[(&str, &str)] = &[
("bootstrap.servers", "owned by the typed `brokers` field"),
(
"metadata.broker.list",
"librdkafka alias of `bootstrap.servers`, owned by the typed \
`brokers` field",
),
(
"statistics.interval.ms",
"owned by the typed `statistics_interval` field",
),
(
"delivery.timeout.ms",
"owned by the typed `delivery_timeout` field, which also bounds how \
long a batch write awaits its delivery reports",
),
(
"message.timeout.ms",
"librdkafka alias of `delivery.timeout.ms`, owned by the typed \
`delivery_timeout` field",
),
(
"message.max.bytes",
"owned by the typed `max_message_bytes` field, which keeps the \
client-side limit aligned with the sink's encode-time guard",
),
(
"acks",
"forced to `all`: the framework commits source offsets once a \
delivery report confirms, so a report under weaker acks would \
turn at-least-once into at-most-once",
),
(
"request.required.acks",
"librdkafka alias of `acks`, which is forced to `all` (a confirmed \
delivery report must mean a durable write)",
),
(
"enable.idempotence",
"forced on: librdkafka's internal retries must not reorder or \
duplicate within a session, and disabling it silently weakens the \
delivery guarantees this sink documents",
),
(
"transactional.id",
"transactions/exactly-once are not supported; the sink's ack model \
is per-batch delivery confirmation, not a two-phase commit",
),
(
"delivery.report.only.error",
"the sink counts every delivery report to acknowledge a batch; \
suppressing success reports would hang every write",
),
(
"enable.gapless.guarantee",
"raises librdkafka fatal errors on any gap, conflicting with the \
framework's retry-and-replay model",
),
];
fn default_shards() -> usize {
1
}
fn default_delivery_timeout() -> Duration {
Duration::from_secs(30)
}
fn default_max_message_bytes() -> ByteSize {
ByteSize(DEFAULT_MAX_MESSAGE_BYTES as u64)
}
fn default_statistics_interval() -> Duration {
Duration::from_secs(5)
}
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum Compression {
None,
Gzip,
Snappy,
Lz4,
Zstd,
}
impl Compression {
fn codec(self) -> &'static str {
match self {
Compression::None => "none",
Compression::Gzip => "gzip",
Compression::Snappy => "snappy",
Compression::Lz4 => "lz4",
Compression::Zstd => "zstd",
}
}
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(deny_unknown_fields)]
pub struct KafkaSinkConfig {
pub brokers: String,
pub topic: String,
#[serde(default = "default_shards")]
pub shards: usize,
#[serde(with = "humantime_serde", default = "default_delivery_timeout")]
pub delivery_timeout: Duration,
#[serde(default = "default_max_message_bytes")]
pub max_message_bytes: ByteSize,
#[serde(with = "humantime_serde", default = "default_statistics_interval")]
pub statistics_interval: Duration,
#[serde(default)]
pub compression: Option<Compression>,
#[serde(default)]
pub batch: BatchConfig,
#[serde(default)]
pub inflight: InflightConfig,
#[serde(default)]
pub retry: RetryConfig,
#[serde(default)]
pub breaker: BreakerConfig,
#[serde(default)]
pub rdkafka: BTreeMap<String, String>,
}
impl KafkaSinkConfig {
pub fn validate(&self) -> Result<(), ConfigError> {
let fail = |msg: String| Err(ConfigError::Validation(format!("sink.kafka: {msg}")));
if self.brokers.trim().is_empty() {
return fail("brokers must not be empty".into());
}
if self.topic.trim().is_empty() {
return fail("topic must not be empty".into());
}
if self.shards == 0 {
return fail("shards must be at least 1".into());
}
if self.delivery_timeout.is_zero() {
return fail(
"delivery_timeout must be positive (librdkafka treats 0 as \
infinite, which would let a batch write hang past every \
framework deadline)"
.into(),
);
}
if self.delivery_timeout.as_millis() > i32::MAX as u128 {
return fail(format!(
"delivery_timeout {:?} exceeds librdkafka's maximum (~24.8 days)",
self.delivery_timeout
));
}
let mmb = self.max_message_bytes.as_u64();
if !(1_000..=1_000_000_000).contains(&mmb) {
return fail(format!(
"max_message_bytes must be in [1000, 1000000000] \
(librdkafka's `message.max.bytes` range), got {mmb}"
));
}
if self.batch.max_rows == 0 || self.batch.max_bytes == 0 {
return fail("batch thresholds must be non-zero".into());
}
if self.inflight.max_per_shard == 0 {
return fail("inflight.max_per_shard must be at least 1".into());
}
if let Err(why) = self.retry.validate() {
return fail(why.to_string());
}
if self.breaker.failure_threshold == 0 {
return fail("breaker.failure_threshold must be at least 1".into());
}
if self.breaker.half_open_probes == 0 {
return fail("breaker.half_open_probes must be at least 1".into());
}
for (key, why) in DENYLIST {
if self.rdkafka.contains_key(*key) {
return fail(format!("rdkafka.\"{key}\" cannot be overridden: {why}"));
}
}
crate::security::check_tls_feature(&self.rdkafka, "sink.kafka")?;
if self.compression.is_some() {
for key in ["compression.codec", "compression.type"] {
if self.rdkafka.contains_key(key) {
return fail(format!(
"rdkafka.\"{key}\" conflicts with the typed \
`compression` field; set one, not both"
));
}
}
}
Ok(())
}
pub(crate) fn client_config(&self) -> rdkafka::ClientConfig {
self.client_config_impl(true)
}
fn probe_client_config(&self) -> rdkafka::ClientConfig {
self.client_config_impl(false)
}
fn client_config_impl(&self, with_statistics: bool) -> rdkafka::ClientConfig {
let mut cc = rdkafka::ClientConfig::new();
for (k, v) in &self.rdkafka {
cc.set(k, v);
}
if let Some(compression) = self.compression {
cc.set("compression.codec", compression.codec());
}
cc.set("bootstrap.servers", &self.brokers);
cc.set("enable.idempotence", "true");
cc.set("acks", "all");
cc.set(
"message.timeout.ms",
self.delivery_timeout.as_millis().to_string(),
);
cc.set(
"message.max.bytes",
self.max_message_bytes.as_u64().to_string(),
);
if with_statistics && !self.statistics_interval.is_zero() {
cc.set(
"statistics.interval.ms",
self.statistics_interval.as_millis().to_string(),
);
}
cc
}
}
#[derive(Debug)]
pub struct KafkaSink {
pub writer: KafkaWriter,
pub endpoints: Vec<Vec<KafkaEndpoint>>,
pub pool: SinkPoolConfig,
probe_endpoints: Arc<Vec<Vec<KafkaEndpoint>>>,
max_message_bytes: usize,
}
impl KafkaSink {
#[must_use]
pub fn encoder_bytes(&self) -> KafkaEncoder<Owned<Vec<u8>>, KafkaBytesEncoder> {
self.encoder_with(KafkaBytesEncoder::new())
}
#[must_use]
pub fn encoder_json<F>(&self) -> KafkaEncoder<F, KafkaJsonEncoder<F>>
where
F: RecFamily,
for<'b> F::Rec<'b>: serde::Serialize,
{
self.encoder_with(KafkaJsonEncoder::new())
}
#[must_use]
pub fn encoder_with<F: RecFamily, M: MessageEncoder<F>>(&self, inner: M) -> KafkaEncoder<F, M> {
KafkaEncoder::with_max_message_bytes(inner, self.max_message_bytes)
}
#[must_use]
pub fn probe_fn(&self) -> SinkProbeFn {
endpoint_probe(self.writer.clone(), Arc::clone(&self.probe_endpoints))
}
}
impl SinkBundle for KafkaSink {
type Writer = KafkaWriter;
fn into_parts(self) -> SinkParts<KafkaWriter> {
let probe = self.probe_fn();
let replica_labels = self
.endpoints
.iter()
.map(|shard| shard.iter().map(|e| e.label().to_string()).collect())
.collect();
SinkParts::new(self.writer, self.endpoints, self.pool)
.with_component_type("kafka")
.with_replica_labels(replica_labels)
.with_probe(probe)
}
}
pub fn from_component_config(section: &ComponentConfig) -> Result<KafkaSink, ConfigError> {
let cfg: KafkaSinkConfig = section.deserialize_into()?;
build(cfg)
}
pub fn build(cfg: KafkaSinkConfig) -> Result<KafkaSink, ConfigError> {
cfg.validate()?;
let stats_slot = Arc::new(Mutex::new(None));
let producer: ThreadedProducer<SinkContext> = cfg
.client_config()
.create_with_context(SinkContext::new(Arc::clone(&stats_slot)))
.map_err(|e| {
ConfigError::Validation(format!("sink.kafka: producer creation failed: {e}"))
})?;
let label = format!("{}/{}", cfg.brokers, cfg.topic);
let endpoint = KafkaEndpoint::new(producer, label.clone());
let endpoints: Vec<Vec<KafkaEndpoint>> =
(0..cfg.shards).map(|_| vec![endpoint.clone()]).collect();
let probe_producer: ThreadedProducer<SinkContext> = cfg
.probe_client_config()
.create_with_context(SinkContext::detached())
.map_err(|e| {
ConfigError::Validation(format!("sink.kafka: probe producer creation failed: {e}"))
})?;
let probe_endpoints = Arc::new(vec![vec![KafkaEndpoint::new(probe_producer, label)]]);
let pool = SinkPoolConfig {
batch: cfg.batch,
inflight: cfg.inflight,
retry: cfg.retry,
breaker: cfg.breaker,
};
Ok(KafkaSink {
writer: KafkaWriter::new(
cfg.topic,
cfg.delivery_timeout,
stats_slot,
!cfg.statistics_interval.is_zero(),
),
endpoints,
pool,
probe_endpoints,
max_message_bytes: cfg.max_message_bytes.as_u64() as usize,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn section(body: &str) -> ComponentConfig {
let yaml = format!("kafka:\n{body}");
let value: serde_yaml::Value = serde_yaml::from_str(&yaml).unwrap();
ComponentConfig::new("kafka", value["kafka"].clone())
}
fn minimal() -> String {
" brokers: localhost:9092\n topic: orders\n".to_string()
}
fn parse(body: &str) -> Result<KafkaSinkConfig, ConfigError> {
let cfg: KafkaSinkConfig = section(body).deserialize_into()?;
cfg.validate()?;
Ok(cfg)
}
#[test]
fn minimal_config_gets_documented_defaults() {
let cfg = parse(&minimal()).unwrap();
assert_eq!(cfg.shards, 1);
assert_eq!(cfg.delivery_timeout, Duration::from_secs(30));
assert_eq!(cfg.max_message_bytes.as_u64(), 1_000_000);
assert_eq!(cfg.statistics_interval, Duration::from_secs(5));
assert_eq!(cfg.compression, None);
assert!(cfg.rdkafka.is_empty());
assert_eq!(cfg.batch, BatchConfig::default());
assert_eq!(cfg.retry, RetryConfig::default());
}
#[test]
fn denylisted_properties_are_rejected_with_reasons() {
for (key, _) in DENYLIST {
let body = format!("{} rdkafka:\n \"{key}\": \"x\"\n", minimal());
let err = parse(&body).unwrap_err();
let msg = err.to_string();
assert!(msg.contains(key), "error names the key: {msg}");
assert!(msg.contains("sink.kafka"), "error names the section: {msg}");
}
}
#[test]
fn delivery_timeout_aliases_are_both_denied() {
for key in ["delivery.timeout.ms", "message.timeout.ms"] {
let body = format!("{} rdkafka:\n \"{key}\": \"1\"\n", minimal());
let msg = parse(&body).unwrap_err().to_string();
assert!(
msg.contains("delivery_timeout"),
"explains ownership: {msg}"
);
}
}
#[test]
fn acks_aliases_are_both_denied() {
for key in ["acks", "request.required.acks"] {
let body = format!("{} rdkafka:\n \"{key}\": \"0\"\n", minimal());
let msg = parse(&body).unwrap_err().to_string();
assert!(msg.contains(key), "error names the key: {msg}");
}
}
#[test]
fn forced_properties_win_over_passthrough() {
let body = format!(
"{} rdkafka:\n linger.ms: \"20\"\n batch.num.messages: \"5000\"\n",
minimal()
);
let cfg = parse(&body).unwrap();
let cc = cfg.client_config();
assert_eq!(cc.get("linger.ms"), Some("20"), "passthrough applies");
assert_eq!(cc.get("batch.num.messages"), Some("5000"));
assert_eq!(cc.get("bootstrap.servers"), Some("localhost:9092"));
assert_eq!(cc.get("enable.idempotence"), Some("true"));
assert_eq!(cc.get("acks"), Some("all"));
assert_eq!(cc.get("message.timeout.ms"), Some("30000"));
assert_eq!(cc.get("message.max.bytes"), Some("1000000"));
assert_eq!(cc.get("statistics.interval.ms"), Some("5000"));
}
#[test]
fn tls_config_matches_build_capability() {
for sec in [
" security.protocol: ssl\n",
" security.protocol: sasl_ssl\n sasl.mechanism: SCRAM-SHA-256\n \
sasl.username: svc\n sasl.password: secret\n",
] {
let body = format!("{} rdkafka:\n{sec}", minimal());
let parsed = parse(&body);
if cfg!(feature = "tls") {
let cfg = parsed.expect("tls build accepts a security config");
let cc = cfg.client_config();
assert!(
cc.get("security.protocol").is_some(),
"security passthrough survives into the client (not denylisted)"
);
let producer: rdkafka::producer::BaseProducer = cc
.create()
.expect("SSL/SASL compiled in: producer creation succeeds");
drop(producer);
} else {
let msg = parsed
.expect_err("non-tls build rejects a security config")
.to_string();
assert!(msg.contains("kafka-tls"), "actionable: {msg}");
}
}
}
#[test]
fn statistics_can_be_disabled_and_probe_config_never_emits() {
let body = format!("{} statistics_interval: 0s\n", minimal());
let cfg = parse(&body).unwrap();
assert_eq!(cfg.client_config().get("statistics.interval.ms"), None);
let enabled = parse(&minimal()).unwrap();
assert_eq!(
enabled.probe_client_config().get("statistics.interval.ms"),
None,
"probe producers never emit statistics"
);
}
#[test]
fn compression_field_maps_and_conflicts_with_passthrough() {
let body = format!("{} compression: lz4\n", minimal());
let cfg = parse(&body).unwrap();
assert_eq!(cfg.client_config().get("compression.codec"), Some("lz4"));
for key in ["compression.codec", "compression.type"] {
let body = format!(
"{} compression: lz4\n rdkafka:\n \"{key}\": \"gzip\"\n",
minimal()
);
let msg = parse(&body).unwrap_err().to_string();
assert!(msg.contains("conflicts"), "conflict is explained: {msg}");
}
let body = format!("{} rdkafka:\n compression.codec: \"gzip\"\n", minimal());
assert!(parse(&body).is_ok());
}
#[test]
fn zero_shards_rejected() {
let body = format!("{} shards: 0\n", minimal());
assert!(parse(&body).unwrap_err().to_string().contains("shards"));
}
#[test]
fn out_of_range_knobs_are_rejected() {
for (body, needle) in [
(
format!("{} delivery_timeout: 0s\n", minimal()),
"delivery_timeout",
),
(
format!("{} max_message_bytes: 10\n", minimal()),
"max_message_bytes",
),
(
format!("{} max_message_bytes: 2GB\n", minimal()),
"max_message_bytes",
),
(
format!("{} retry:\n multiplier: 0.5\n", minimal()),
"retry.multiplier",
),
(
format!("{} retry:\n jitter: 1.5\n", minimal()),
"retry.jitter",
),
(
format!("{} breaker:\n failure_threshold: 0\n", minimal()),
"breaker.failure_threshold",
),
(
format!("{} inflight:\n max_per_shard: 0\n", minimal()),
"inflight.max_per_shard",
),
] {
let msg = parse(&body).unwrap_err().to_string();
assert!(msg.contains(needle), "`{needle}` in: {msg}");
}
}
#[test]
fn tuning_sections_map_to_pool_config() {
let body = format!(
"{} shards: 3\n batch:\n max_rows: 1000\n max_bytes: 1MiB\n linger: 250ms\n \
inflight:\n max_per_shard: 4\n retry:\n initial: 50ms\n max: 5s\n \
breaker:\n failure_threshold: 7\n",
minimal()
);
let sink = build(parse(&body).unwrap()).unwrap();
assert_eq!(sink.endpoints.len(), 3, "one worker per shard");
assert!(
sink.endpoints.iter().all(|shard| shard.len() == 1),
"single-replica shards"
);
assert_eq!(sink.pool.batch.max_rows, 1000);
assert_eq!(sink.pool.batch.max_bytes, 1024 * 1024);
assert_eq!(sink.pool.batch.linger, Duration::from_millis(250));
assert_eq!(sink.pool.inflight.max_per_shard, 4);
assert_eq!(sink.pool.retry.initial, Duration::from_millis(50));
assert_eq!(sink.pool.breaker.failure_threshold, 7);
}
#[test]
fn bundle_exposes_kafka_component_type_and_labels() {
let sink = build(parse(&minimal()).unwrap()).unwrap();
let parts = sink.into_parts();
assert_eq!(parts.component_type, "kafka");
assert_eq!(
parts.replica_labels,
vec![vec!["localhost:9092/orders".to_string()]]
);
assert!(parts.probe.is_some(), "readiness probe attached");
}
#[test]
fn idempotence_incompatible_passthrough_fails_at_build() {
let body = format!("{} rdkafka:\n max.in.flight: \"10\"\n", minimal());
let err = build(parse(&body).unwrap()).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("producer creation failed"),
"surfaced at startup: {msg}"
);
}
#[test]
fn unknown_fields_are_rejected() {
let body = format!("{} topics: [a, b]\n", minimal());
assert!(
section(&body)
.deserialize_into::<KafkaSinkConfig>()
.is_err()
);
}
#[test]
fn empty_required_fields_error_clearly() {
for body in [
" brokers: \"\"\n topic: t\n",
" brokers: b\n topic: \"\"\n",
] {
assert!(parse(body).is_err());
}
}
}