1use anyhow::bail;
13use aws_config::SdkConfig;
14use fancy_regex::Regex;
15use std::collections::{BTreeMap, btree_map};
16use std::error::Error;
17use std::io;
18use std::net::{SocketAddr, ToSocketAddrs};
19use std::str::FromStr;
20use std::sync::Arc;
21use std::sync::Mutex;
22use std::time::Duration;
23use tokio::sync::watch;
24
25use anyhow::{Context, anyhow};
26use crossbeam::channel::{Receiver, Sender, unbounded};
27use mz_ore::collections::CollectionExt;
28use mz_ore::error::ErrorExt;
29use mz_ore::future::InTask;
30use mz_ssh_util::tunnel::{SshTimeoutConfig, SshTunnelConfig, SshTunnelStatus};
31use mz_ssh_util::tunnel_manager::{ManagedSshTunnelHandle, SshTunnelManager};
32use rdkafka::client::{Client, NativeClient, OAuthToken};
33use rdkafka::config::{
34 ClientConfig, FromClientConfig, FromClientConfigAndContext, RDKafkaLogLevel,
35};
36use rdkafka::consumer::{ConsumerContext, Rebalance};
37use rdkafka::error::{KafkaError, KafkaResult, RDKafkaErrorCode};
38use rdkafka::producer::{DefaultProducerContext, DeliveryResult, ProducerContext};
39use rdkafka::types::RDKafkaRespErr;
40use rdkafka::util::Timeout;
41use rdkafka::{ClientContext, Statistics, TopicPartitionList};
42use serde::{Deserialize, Serialize};
43use tokio::runtime::Handle;
44use tracing::{Level, debug, info, trace, warn};
45
46use crate::aws;
47
48pub const DEFAULT_TOPIC_METADATA_REFRESH_INTERVAL: Duration = Duration::from_secs(30);
54
55pub struct MzClientContext {
62 error_tx: Sender<MzKafkaError>,
64 statistics_tx: watch::Sender<Statistics>,
67}
68
69impl Default for MzClientContext {
70 fn default() -> Self {
71 Self::with_errors().0
72 }
73}
74
75impl MzClientContext {
76 pub fn with_errors() -> (Self, Receiver<MzKafkaError>) {
81 let (error_tx, error_rx) = unbounded();
82 let (statistics_tx, _) = watch::channel(Default::default());
83 let ctx = Self {
84 error_tx,
85 statistics_tx,
86 };
87 (ctx, error_rx)
88 }
89
90 pub fn subscribe_statistics(&self) -> watch::Receiver<Statistics> {
93 self.statistics_tx.subscribe()
94 }
95
96 fn record_error(&self, msg: &str) {
97 let err = match MzKafkaError::from_str(msg) {
98 Ok(err) => err,
99 Err(()) => {
100 warn!(original_error = msg, "failed to parse kafka error");
101 MzKafkaError::Internal(msg.to_owned())
102 }
103 };
104 let _ = self.error_tx.send(err);
106 }
107}
108
109#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
111pub enum MzKafkaError {
112 #[error("Invalid username or password")]
114 InvalidCredentials,
115 #[error("Invalid CA certificate")]
117 InvalidCACertificate,
118 #[error("Disconnected during handshake; broker might require SSL encryption")]
120 SSLEncryptionMaybeRequired,
121 #[error("Broker does not support SSL connections")]
123 SSLUnsupported,
124 #[error("Broker did not provide a certificate")]
126 BrokerCertificateMissing,
127 #[error("Failed to verify broker certificate")]
129 InvalidBrokerCertificate,
130 #[error("Connection reset: {0}")]
132 ConnectionReset(String),
133 #[error("Connection timeout")]
135 ConnectionTimeout,
136 #[error("Failed to resolve hostname")]
138 HostnameResolutionFailed,
139 #[error("Unsupported SASL mechanism")]
141 UnsupportedSASLMechanism,
142 #[error("Unsupported broker version")]
144 UnsupportedBrokerVersion,
145 #[error("Broker transport failure")]
147 BrokerTransportFailure,
148 #[error("All brokers down")]
150 AllBrokersDown,
151 #[error("SASL authentication required")]
153 SaslAuthenticationRequired,
154 #[error("SASL authentication failed")]
156 SaslAuthenticationFailed,
157 #[error("SSL authentication required")]
159 SslAuthenticationRequired,
160 #[error("Unknown topic or partition")]
162 UnknownTopicOrPartition,
163 #[error("Internal kafka error: {0}")]
165 Internal(String),
166}
167
168impl FromStr for MzKafkaError {
169 type Err = ();
170
171 fn from_str(s: &str) -> Result<Self, Self::Err> {
172 if s.contains("Authentication failed: Invalid username or password") {
173 Ok(Self::InvalidCredentials)
174 } else if s.contains("broker certificate could not be verified") {
175 Ok(Self::InvalidCACertificate)
176 } else if s.contains("connecting to a SSL listener?") {
177 Ok(Self::SSLEncryptionMaybeRequired)
178 } else if s.contains("client SSL authentication might be required") {
179 Ok(Self::SslAuthenticationRequired)
180 } else if s.contains("connecting to a PLAINTEXT broker listener") {
181 Ok(Self::SSLUnsupported)
182 } else if s.contains("Broker did not provide a certificate") {
183 Ok(Self::BrokerCertificateMissing)
184 } else if s.contains("Failed to verify broker certificate: ") {
185 Ok(Self::InvalidBrokerCertificate)
186 } else if let Some((_prefix, inner)) = s.split_once("Send failed: ") {
187 Ok(Self::ConnectionReset(inner.to_owned()))
188 } else if let Some((_prefix, inner)) = s.split_once("Receive failed: ") {
189 Ok(Self::ConnectionReset(inner.to_owned()))
190 } else if s.contains("request(s) timed out: disconnect") {
191 Ok(Self::ConnectionTimeout)
192 } else if s.contains("Failed to resolve") {
193 Ok(Self::HostnameResolutionFailed)
194 } else if s.contains("mechanism handshake failed:") {
195 Ok(Self::UnsupportedSASLMechanism)
196 } else if s.contains(
197 "verify that security.protocol is correctly configured, \
198 broker might require SASL authentication",
199 ) {
200 Ok(Self::SaslAuthenticationRequired)
201 } else if s.contains("SASL authentication error: Authentication failed") {
202 Ok(Self::SaslAuthenticationFailed)
203 } else if s
204 .contains("incorrect security.protocol configuration (connecting to a SSL listener?)")
205 {
206 Ok(Self::SslAuthenticationRequired)
207 } else if s.contains("probably due to broker version < 0.10") {
208 Ok(Self::UnsupportedBrokerVersion)
209 } else if s.contains("Disconnected while requesting ApiVersion")
210 || s.contains("Broker transport failure")
211 || s.contains("Connection refused")
212 {
213 Ok(Self::BrokerTransportFailure)
214 } else if Regex::new(r"(\d+)/\1 brokers are down")
215 .unwrap()
216 .is_match(s)
217 .unwrap_or_default()
218 {
219 Ok(Self::AllBrokersDown)
220 } else if s.contains("Unknown topic or partition") || s.contains("Unknown partition") {
221 Ok(Self::UnknownTopicOrPartition)
222 } else {
223 Err(())
224 }
225 }
226}
227
228impl ClientContext for MzClientContext {
229 fn log(&self, level: rdkafka::config::RDKafkaLogLevel, fac: &str, log_message: &str) {
230 use rdkafka::config::RDKafkaLogLevel::*;
231
232 if matches!(level, Emerg | Alert | Critical | Error) || fac == "FAIL" {
240 self.record_error(log_message);
241 }
242
243 match level {
246 Emerg | Alert | Critical | Error => {
247 warn!(target: "librdkafka", "error: {} {}", fac, log_message);
252 }
253 Warning => warn!(target: "librdkafka", "warning: {} {}", fac, log_message),
254 Notice => info!(target: "librdkafka", "{} {}", fac, log_message),
255 Info => info!(target: "librdkafka", "{} {}", fac, log_message),
256 Debug => debug!(target: "librdkafka", "{} {}", fac, log_message),
257 }
258 }
259
260 fn stats(&self, statistics: Statistics) {
261 self.statistics_tx.send_replace(statistics);
262 }
263
264 fn error(&self, error: KafkaError, reason: &str) {
265 self.record_error(reason);
266 warn!(target: "librdkafka", "error: {}: {}", error, reason);
268 }
269}
270
271impl ConsumerContext for MzClientContext {}
272
273impl ProducerContext for MzClientContext {
274 type DeliveryOpaque = <DefaultProducerContext as ProducerContext>::DeliveryOpaque;
275 fn delivery(
276 &self,
277 delivery_result: &DeliveryResult<'_>,
278 delivery_opaque: Self::DeliveryOpaque,
279 ) {
280 DefaultProducerContext.delivery(delivery_result, delivery_opaque);
281 }
282}
283
284#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd)]
286pub struct BrokerAddr {
287 pub host: String,
289 pub port: u16,
291}
292
293impl BrokerAddr {
294 pub fn to_socket_addrs(&self) -> Result<Vec<SocketAddr>, io::Error> {
296 Ok((self.host.as_str(), self.port).to_socket_addrs()?.collect())
297 }
298}
299
300#[derive(Debug, Clone)]
304pub struct BrokerRewrite {
305 pub host: String,
307 pub port: Option<u16>,
311}
312
313impl BrokerRewrite {
314 pub fn rewrite(&self, address: &BrokerAddr) -> BrokerAddr {
316 BrokerAddr {
317 host: self.host.clone(),
318 port: self.port.unwrap_or(address.port),
319 }
320 }
321}
322
323#[derive(Clone)]
324enum BrokerRewriteHandle {
325 Simple(BrokerRewrite),
326 SshTunnel(
327 ManagedSshTunnelHandle,
329 ),
330 FailedDefaultSshTunnel(String),
333}
334
335#[derive(Clone)]
336pub struct ConnectionRulePattern {
338 pub prefix_wildcard: bool,
340 pub literal_match: String,
342 pub suffix_wildcard: bool,
344}
345
346impl std::fmt::Display for ConnectionRulePattern {
347 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
348 if self.prefix_wildcard {
349 f.write_str("*")?;
350 }
351 f.write_str(&self.literal_match)?;
352 if self.suffix_wildcard {
353 f.write_str("*")?;
354 }
355 Ok(())
356 }
357}
358
359impl ConnectionRulePattern {
360 pub fn matches(&self, address: &str) -> bool {
362 if self.prefix_wildcard {
363 if self.suffix_wildcard {
364 address.contains(&self.literal_match)
365 } else {
366 address.ends_with(&self.literal_match)
367 }
368 } else if self.suffix_wildcard {
369 address.starts_with(&self.literal_match)
370 } else {
371 address == self.literal_match
372 }
373 }
374}
375
376#[derive(Clone)]
377pub struct HostMappingRules {
379 pub rules: Vec<(ConnectionRulePattern, BrokerRewrite)>,
381}
382
383impl HostMappingRules {
384 pub fn rewrite(&self, src: &BrokerAddr) -> Option<BrokerAddr> {
387 let address = format!("{}:{}", src.host, src.port);
388 for (pattern, dst) in &self.rules {
389 if pattern.matches(&address) {
390 let result = dst.rewrite(src);
391 info!(
392 "HostMappingRules: broker {}:{} matched pattern '{}' -> rewriting to {}:{}",
393 src.host, src.port, pattern, result.host, result.port,
394 );
395 return Some(result);
396 }
397 }
398
399 warn!(
400 "HostMappingRules: broker {}:{} matched no rules, using original address",
401 src.host, src.port,
402 );
403 None
404 }
405}
406
407#[derive(Clone)]
410pub enum TunnelConfig {
411 Ssh(SshTunnelConfig),
413 StaticHost(String),
415 Rules(HostMappingRules),
417 None,
419}
420
421#[derive(Clone)]
423pub struct TunnelingClientContext<C> {
424 inner: C,
425 rewrites: Arc<Mutex<BTreeMap<BrokerAddr, BrokerRewriteHandle>>>,
426 default_tunnel: TunnelConfig,
427 in_task: InTask,
428 ssh_tunnel_manager: SshTunnelManager,
429 ssh_timeout_config: SshTimeoutConfig,
430 aws_config: Option<SdkConfig>,
431 runtime: Handle,
432}
433
434impl<C> TunnelingClientContext<C> {
435 pub fn new(
437 inner: C,
438 runtime: Handle,
439 ssh_tunnel_manager: SshTunnelManager,
440 ssh_timeout_config: SshTimeoutConfig,
441 aws_config: Option<SdkConfig>,
442 in_task: InTask,
443 ) -> TunnelingClientContext<C> {
444 TunnelingClientContext {
445 inner,
446 rewrites: Arc::new(Mutex::new(BTreeMap::new())),
447 default_tunnel: TunnelConfig::None,
448 in_task,
449 ssh_tunnel_manager,
450 ssh_timeout_config,
451 aws_config,
452 runtime,
453 }
454 }
455
456 pub fn set_default_tunnel(&mut self, tunnel: TunnelConfig) {
461 self.default_tunnel = tunnel;
462 }
463
464 pub async fn add_ssh_tunnel(
471 &self,
472 broker: BrokerAddr,
473 tunnel: SshTunnelConfig,
474 ) -> Result<(), anyhow::Error> {
475 let ssh_tunnel = self
476 .ssh_tunnel_manager
477 .connect(
478 tunnel,
479 &broker.host,
480 broker.port,
481 self.ssh_timeout_config,
482 self.in_task,
483 )
484 .await
485 .context("creating ssh tunnel")?;
486
487 let mut rewrites = self.rewrites.lock().expect("poisoned");
488 rewrites.insert(broker, BrokerRewriteHandle::SshTunnel(ssh_tunnel));
489 Ok(())
490 }
491
492 pub fn add_broker_rewrite(&self, broker: BrokerAddr, rewrite: BrokerRewrite) {
498 let mut rewrites = self.rewrites.lock().expect("poisoned");
499 rewrites.insert(broker, BrokerRewriteHandle::Simple(rewrite));
500 }
501
502 pub fn inner(&self) -> &C {
504 &self.inner
505 }
506
507 pub fn tunnel_status(&self) -> SshTunnelStatus {
510 self.rewrites
511 .lock()
512 .expect("poisoned")
513 .values()
514 .map(|handle| match handle {
515 BrokerRewriteHandle::SshTunnel(s) => s.check_status(),
516 BrokerRewriteHandle::FailedDefaultSshTunnel(e) => {
517 SshTunnelStatus::Errored(e.clone())
518 }
519 BrokerRewriteHandle::Simple(_) => SshTunnelStatus::Running,
520 })
521 .fold(SshTunnelStatus::Running, |acc, status| {
522 match (acc, status) {
523 (SshTunnelStatus::Running, SshTunnelStatus::Errored(e))
524 | (SshTunnelStatus::Errored(e), SshTunnelStatus::Running) => {
525 SshTunnelStatus::Errored(e)
526 }
527 (SshTunnelStatus::Errored(err), SshTunnelStatus::Errored(e)) => {
528 SshTunnelStatus::Errored(format!("{}, {}", err, e))
529 }
530 (SshTunnelStatus::Running, SshTunnelStatus::Running) => {
531 SshTunnelStatus::Running
532 }
533 }
534 })
535 }
536}
537
538impl<C> ClientContext for TunnelingClientContext<C>
539where
540 C: ClientContext,
541{
542 const ENABLE_REFRESH_OAUTH_TOKEN: bool = true;
543
544 fn generate_oauth_token(
545 &self,
546 _oauthbearer_config: Option<&str>,
547 ) -> Result<OAuthToken, Box<dyn Error>> {
548 info!(target: "librdkafka", "generating OAuth token");
559
560 let generate = || {
561 let Some(sdk_config) = &self.aws_config else {
562 bail!("internal error: AWS configuration missing");
563 };
564
565 self.runtime.block_on(aws::generate_auth_token(sdk_config))
566 };
567
568 match generate() {
569 Ok((token, lifetime_ms)) => {
570 info!(target: "librdkafka", %lifetime_ms, "successfully generated OAuth token");
571 trace!(target: "librdkafka", %token);
572 Ok(OAuthToken {
573 token,
574 lifetime_ms,
575 principal_name: "".to_string(),
576 })
577 }
578 Err(e) => {
579 warn!(target: "librdkafka", "failed to generate OAuth token: {e:#}");
580 Err(e.into())
581 }
582 }
583 }
584
585 fn resolve_broker_addr(&self, host: &str, port: u16) -> Result<Vec<SocketAddr>, io::Error> {
590 info!("kafka: resolve_broker_addr called for {}:{}", host, port);
591 let return_rewrite = |rewrite: &BrokerRewriteHandle| -> Result<Vec<SocketAddr>, io::Error> {
592 let rewrite = match rewrite {
593 BrokerRewriteHandle::Simple(rewrite) => rewrite.clone(),
594 BrokerRewriteHandle::SshTunnel(ssh_tunnel) => {
595 let addr = ssh_tunnel.local_addr();
598 BrokerRewrite {
599 host: addr.ip().to_string(),
600 port: Some(addr.port()),
601 }
602 }
603 BrokerRewriteHandle::FailedDefaultSshTunnel(_) => {
604 unreachable!()
605 }
606 };
607 let rewrite_port = rewrite.port.unwrap_or(port);
608
609 info!(
610 "rewriting broker {}:{} to {}:{}",
611 host, port, rewrite.host, rewrite_port
612 );
613
614 (rewrite.host, rewrite_port)
615 .to_socket_addrs()
616 .map(|addrs| addrs.collect())
617 };
618
619 let addr = BrokerAddr {
620 host: host.into(),
621 port,
622 };
623 let rewrite = self.rewrites.lock().expect("poisoned").get(&addr).cloned();
624
625 match rewrite {
626 None | Some(BrokerRewriteHandle::FailedDefaultSshTunnel(_)) => {
628 match &self.default_tunnel {
630 TunnelConfig::Ssh(default_tunnel) => {
633 let ssh_tunnel = self.runtime.block_on(async {
637 self.ssh_tunnel_manager
638 .connect(
639 default_tunnel.clone(),
640 host,
641 port,
642 self.ssh_timeout_config,
643 self.in_task,
644 )
645 .await
646 });
647 match ssh_tunnel {
648 Ok(ssh_tunnel) => {
650 let mut rewrites = self.rewrites.lock().expect("poisoned");
651 let rewrite = match rewrites.entry(addr.clone()) {
652 btree_map::Entry::Occupied(mut o)
653 if matches!(
654 o.get(),
655 BrokerRewriteHandle::FailedDefaultSshTunnel(_)
656 ) =>
657 {
658 o.insert(BrokerRewriteHandle::SshTunnel(
659 ssh_tunnel.clone(),
660 ));
661 o.into_mut()
662 }
663 btree_map::Entry::Occupied(o) => o.into_mut(),
664 btree_map::Entry::Vacant(v) => {
665 v.insert(BrokerRewriteHandle::SshTunnel(ssh_tunnel.clone()))
666 }
667 };
668
669 return_rewrite(rewrite)
670 }
671 Err(e) => {
673 warn!(
674 "failed to create ssh tunnel for {:?}: {}",
675 addr,
676 e.display_with_causes()
677 );
678
679 let mut rewrites = self.rewrites.lock().expect("poisoned");
681 rewrites.entry(addr.clone()).or_insert_with(|| {
682 BrokerRewriteHandle::FailedDefaultSshTunnel(
683 e.to_string_with_causes(),
684 )
685 });
686
687 Err(io::Error::new(
688 io::ErrorKind::Other,
689 "creating SSH tunnel failed",
690 ))
691 }
692 }
693 }
694 TunnelConfig::StaticHost(host) => (host.as_str(), port)
696 .to_socket_addrs()
697 .map(|addrs| addrs.collect()),
698 TunnelConfig::Rules(rules) => {
700 let resolved = rules.rewrite(&addr).unwrap_or_else(|| addr.clone());
702 match resolved.to_socket_addrs() {
703 Ok(addrs) => {
704 info!(
705 "kafka: resolve_broker_addr {}:{} -> {}:{} resolved to {:?}",
706 host, port, resolved.host, resolved.port, addrs,
707 );
708 Ok(addrs)
709 }
710 Err(e) => {
711 warn!(
712 "kafka: resolve_broker_addr {}:{} -> {}:{} DNS resolution FAILED: {e}",
713 host, port, resolved.host, resolved.port,
714 );
715 Err(e)
716 }
717 }
718 }
719 TunnelConfig::None => {
721 (host, port).to_socket_addrs().map(|addrs| addrs.collect())
722 }
723 }
724 }
725 Some(rewrite) => {
727 info!(
728 "kafka: resolve_broker_addr {}:{} using cached rewrite",
729 host, port
730 );
731 return_rewrite(&rewrite)
732 }
733 }
734 }
735
736 fn log(&self, level: RDKafkaLogLevel, fac: &str, log_message: &str) {
737 self.inner.log(level, fac, log_message)
738 }
739
740 fn error(&self, error: KafkaError, reason: &str) {
741 self.inner.error(error, reason)
742 }
743
744 fn stats(&self, statistics: Statistics) {
745 self.inner.stats(statistics)
746 }
747
748 fn stats_raw(&self, statistics: &[u8]) {
749 self.inner.stats_raw(statistics)
750 }
751}
752
753impl<C> ConsumerContext for TunnelingClientContext<C>
754where
755 C: ConsumerContext,
756{
757 fn rebalance(
758 &self,
759 native_client: &NativeClient,
760 err: RDKafkaRespErr,
761 tpl: &mut TopicPartitionList,
762 ) {
763 self.inner.rebalance(native_client, err, tpl)
764 }
765
766 fn pre_rebalance<'a>(&self, rebalance: &Rebalance<'a>) {
767 self.inner.pre_rebalance(rebalance)
768 }
769
770 fn post_rebalance<'a>(&self, rebalance: &Rebalance<'a>) {
771 self.inner.post_rebalance(rebalance)
772 }
773
774 fn commit_callback(&self, result: KafkaResult<()>, offsets: &TopicPartitionList) {
775 self.inner.commit_callback(result, offsets)
776 }
777
778 fn main_queue_min_poll_interval(&self) -> Timeout {
779 self.inner.main_queue_min_poll_interval()
780 }
781}
782
783impl<C> ProducerContext for TunnelingClientContext<C>
784where
785 C: ProducerContext,
786{
787 type DeliveryOpaque = C::DeliveryOpaque;
788
789 fn delivery(
790 &self,
791 delivery_result: &DeliveryResult<'_>,
792 delivery_opaque: Self::DeliveryOpaque,
793 ) {
794 self.inner.delivery(delivery_result, delivery_opaque)
795 }
796}
797
798pub type PartitionId = i32;
800
801#[derive(Debug, thiserror::Error)]
803pub enum GetPartitionsError {
804 #[error("Topic does not exist")]
806 TopicDoesNotExist,
807 #[error(transparent)]
809 Kafka(#[from] KafkaError),
810 #[error(transparent)]
812 Other(#[from] anyhow::Error),
813}
814
815pub fn get_partitions<C: ClientContext>(
817 client: &Client<C>,
818 topic: &str,
819 timeout: Duration,
820) -> Result<Vec<PartitionId>, GetPartitionsError> {
821 let meta = client.fetch_metadata(Some(topic), timeout)?;
822 if meta.topics().len() != 1 {
823 Err(anyhow!(
824 "topic {} has {} metadata entries; expected 1",
825 topic,
826 meta.topics().len()
827 ))?;
828 }
829
830 fn check_err(err: Option<RDKafkaRespErr>) -> Result<(), GetPartitionsError> {
831 match err.map(RDKafkaErrorCode::from) {
832 Some(RDKafkaErrorCode::UnknownTopic | RDKafkaErrorCode::UnknownTopicOrPartition) => {
833 Err(GetPartitionsError::TopicDoesNotExist)
834 }
835 Some(code) => Err(anyhow!(code))?,
836 None => Ok(()),
837 }
838 }
839
840 let meta_topic = meta.topics().into_element();
841 check_err(meta_topic.error())?;
842
843 if meta_topic.name() != topic {
844 Err(anyhow!(
845 "got results for wrong topic {} (expected {})",
846 meta_topic.name(),
847 topic
848 ))?;
849 }
850
851 let mut partition_ids = Vec::with_capacity(meta_topic.partitions().len());
852 for partition_meta in meta_topic.partitions() {
853 check_err(partition_meta.error())?;
854
855 partition_ids.push(partition_meta.id());
856 }
857
858 if partition_ids.len() == 0 {
859 Err(GetPartitionsError::TopicDoesNotExist)?;
860 }
861
862 Ok(partition_ids)
863}
864
865pub const DEFAULT_KEEPALIVE: bool = true;
867pub const DEFAULT_SOCKET_TIMEOUT: Duration = Duration::from_secs(60);
870pub const DEFAULT_TRANSACTION_TIMEOUT: Duration = Duration::from_secs(600);
873pub const DEFAULT_SOCKET_CONNECTION_SETUP_TIMEOUT: Duration = Duration::from_secs(30);
876pub const DEFAULT_FETCH_METADATA_TIMEOUT: Duration = Duration::from_secs(10);
878pub const DEFAULT_PROGRESS_RECORD_FETCH_TIMEOUT: Duration = Duration::from_secs(90);
881
882#[derive(Copy, Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
884pub struct TimeoutConfig {
885 pub keepalive: bool,
887 pub socket_timeout: Duration,
890 pub transaction_timeout: Duration,
892 pub socket_connection_setup_timeout: Duration,
894 pub fetch_metadata_timeout: Duration,
896 pub progress_record_fetch_timeout: Duration,
898}
899
900impl Default for TimeoutConfig {
901 fn default() -> Self {
902 TimeoutConfig {
903 keepalive: DEFAULT_KEEPALIVE,
904 socket_timeout: DEFAULT_SOCKET_TIMEOUT,
905 transaction_timeout: DEFAULT_TRANSACTION_TIMEOUT,
906 socket_connection_setup_timeout: DEFAULT_SOCKET_CONNECTION_SETUP_TIMEOUT,
907 fetch_metadata_timeout: DEFAULT_FETCH_METADATA_TIMEOUT,
908 progress_record_fetch_timeout: DEFAULT_PROGRESS_RECORD_FETCH_TIMEOUT,
909 }
910 }
911}
912
913impl TimeoutConfig {
914 pub fn build(
917 keepalive: bool,
918 socket_timeout: Option<Duration>,
919 transaction_timeout: Duration,
920 socket_connection_setup_timeout: Duration,
921 fetch_metadata_timeout: Duration,
922 progress_record_fetch_timeout: Option<Duration>,
923 ) -> TimeoutConfig {
924 let transaction_timeout = if transaction_timeout.as_millis() > i32::MAX.try_into().unwrap()
931 {
932 warn!(
933 "transaction_timeout ({transaction_timeout:?}) greater than max \
934 of {}, defaulting to the default of {DEFAULT_TRANSACTION_TIMEOUT:?}",
935 i32::MAX
936 );
937 DEFAULT_TRANSACTION_TIMEOUT
938 } else if transaction_timeout.as_millis() < 1000 {
939 warn!(
940 "transaction_timeout ({transaction_timeout:?}) less than max \
941 of 1000ms, defaulting to the default of {DEFAULT_TRANSACTION_TIMEOUT:?}"
942 );
943 DEFAULT_TRANSACTION_TIMEOUT
944 } else {
945 transaction_timeout
946 };
947
948 let progress_record_fetch_timeout_derived_default =
949 std::cmp::max(transaction_timeout, DEFAULT_PROGRESS_RECORD_FETCH_TIMEOUT);
950 let progress_record_fetch_timeout =
951 progress_record_fetch_timeout.unwrap_or(progress_record_fetch_timeout_derived_default);
952 let progress_record_fetch_timeout = if progress_record_fetch_timeout < transaction_timeout {
953 warn!(
954 "progress record fetch ({progress_record_fetch_timeout:?}) less than transaction \
955 timeout ({transaction_timeout:?}), defaulting to transaction timeout {transaction_timeout:?}",
956 );
957 transaction_timeout
958 } else {
959 progress_record_fetch_timeout
960 };
961
962 let max_socket_timeout = std::cmp::min(
965 transaction_timeout + Duration::from_millis(100),
966 Duration::from_secs(300),
967 );
968 let socket_timeout_derived_default =
969 std::cmp::min(max_socket_timeout, DEFAULT_SOCKET_TIMEOUT);
970 let socket_timeout = socket_timeout.unwrap_or(socket_timeout_derived_default);
971 let socket_timeout = if socket_timeout > max_socket_timeout {
972 warn!(
973 "socket_timeout ({socket_timeout:?}) greater than max \
974 of min(30000, transaction.timeout.ms + 100 ({})), \
975 defaulting to the maximum of {max_socket_timeout:?}",
976 transaction_timeout.as_millis() + 100
977 );
978 max_socket_timeout
979 } else if socket_timeout.as_millis() < 10 {
980 warn!(
981 "socket_timeout ({socket_timeout:?}) less than min \
982 of 10ms, defaulting to the default of {socket_timeout_derived_default:?}"
983 );
984 socket_timeout_derived_default
985 } else {
986 socket_timeout
987 };
988
989 let socket_connection_setup_timeout =
990 if socket_connection_setup_timeout.as_millis() > i32::MAX.try_into().unwrap() {
991 warn!(
992 "socket_connection_setup_timeout ({socket_connection_setup_timeout:?}) \
993 greater than max of {}ms, defaulting to the default \
994 of {DEFAULT_SOCKET_CONNECTION_SETUP_TIMEOUT:?}",
995 i32::MAX,
996 );
997 DEFAULT_SOCKET_CONNECTION_SETUP_TIMEOUT
998 } else if socket_connection_setup_timeout.as_millis() < 10 {
999 warn!(
1000 "socket_connection_setup_timeout ({socket_connection_setup_timeout:?}) \
1001 less than max of 10ms, defaulting to the default of \
1002 {DEFAULT_SOCKET_CONNECTION_SETUP_TIMEOUT:?}"
1003 );
1004 DEFAULT_SOCKET_CONNECTION_SETUP_TIMEOUT
1005 } else {
1006 socket_connection_setup_timeout
1007 };
1008
1009 TimeoutConfig {
1010 keepalive,
1011 socket_timeout,
1012 transaction_timeout,
1013 socket_connection_setup_timeout,
1014 fetch_metadata_timeout,
1015 progress_record_fetch_timeout,
1016 }
1017 }
1018}
1019
1020pub fn create_new_client_config_simple() -> ClientConfig {
1023 create_new_client_config(tracing::Level::INFO, Default::default())
1024}
1025
1026pub fn create_new_client_config(
1030 tracing_level: Level,
1031 timeout_config: TimeoutConfig,
1032) -> ClientConfig {
1033 #[allow(clippy::disallowed_methods)]
1034 let mut config = ClientConfig::new();
1035
1036 let level = if tracing_level >= Level::DEBUG {
1037 RDKafkaLogLevel::Debug
1038 } else if tracing_level >= Level::INFO {
1039 RDKafkaLogLevel::Info
1040 } else if tracing_level >= Level::WARN {
1041 RDKafkaLogLevel::Warning
1042 } else {
1043 RDKafkaLogLevel::Error
1044 };
1045 tracing::debug!(target: "librdkafka", level = ?level, "Determined log level for librdkafka");
1062 config.set_log_level(level);
1063
1064 if tracing_level >= Level::DEBUG {
1069 tracing::debug!(target: "librdkafka", "Enabling debug logs for rdkafka");
1070 config.set("debug", "all");
1071 }
1072
1073 if timeout_config.keepalive {
1074 config.set("socket.keepalive.enable", "true");
1075 }
1076
1077 config.set(
1078 "socket.timeout.ms",
1079 timeout_config.socket_timeout.as_millis().to_string(),
1080 );
1081 config.set(
1082 "socket.connection.setup.timeout.ms",
1083 timeout_config
1084 .socket_connection_setup_timeout
1085 .as_millis()
1086 .to_string(),
1087 );
1088
1089 config
1090}
1091
1092pub fn create<T: FromClientConfig>(config: &ClientConfig) -> KafkaResult<T> {
1097 #[allow(clippy::disallowed_methods)]
1098 let client = config.create();
1099 drop(openssl::error::ErrorStack::get());
1100 client
1101}
1102
1103pub fn create_with_context<C, T>(config: &ClientConfig, context: C) -> KafkaResult<T>
1107where
1108 C: ClientContext,
1109 T: FromClientConfigAndContext<C>,
1110{
1111 #[allow(clippy::disallowed_methods)]
1112 let client = config.create_with_context(context);
1113 drop(openssl::error::ErrorStack::get());
1122 client
1123}
1124
1125#[cfg(test)]
1126mod tests {
1127 use super::*;
1128
1129 use openssl::asn1::Asn1Time;
1130 use openssl::hash::MessageDigest;
1131 use openssl::pkey::PKey;
1132 use openssl::rsa::Rsa;
1133 use openssl::x509::{X509, X509NameBuilder};
1134 use rdkafka::consumer::BaseConsumer;
1135
1136 fn self_signed_cert_pem() -> String {
1137 let key = PKey::from_rsa(Rsa::generate(2048).unwrap()).unwrap();
1138 let mut name = X509NameBuilder::new().unwrap();
1139 name.append_entry_by_text("CN", "test").unwrap();
1140 let name = name.build();
1141 let mut cert = X509::builder().unwrap();
1142 cert.set_version(2).unwrap();
1143 cert.set_subject_name(&name).unwrap();
1144 cert.set_issuer_name(&name).unwrap();
1145 cert.set_pubkey(&key).unwrap();
1146 cert.set_not_before(&Asn1Time::days_from_now(0).unwrap())
1147 .unwrap();
1148 cert.set_not_after(&Asn1Time::days_from_now(1).unwrap())
1149 .unwrap();
1150 cert.sign(&key, MessageDigest::sha256()).unwrap();
1151 String::from_utf8(cert.build().to_pem().unwrap()).unwrap()
1152 }
1153
1154 #[mz_ore::test]
1155 #[cfg_attr(miri, ignore)] fn test_create_with_context_clears_openssl_error_queue() {
1157 let mut config = create_new_client_config_simple();
1158 config.set("bootstrap.servers", "localhost:1");
1160 config.set("security.protocol", "ssl");
1161 config.set("ssl.ca.pem", self_signed_cert_pem());
1162 let _consumer: BaseConsumer =
1163 create_with_context(&config, rdkafka::consumer::DefaultConsumerContext).unwrap();
1164 let errors = openssl::error::ErrorStack::get();
1165 assert!(
1166 errors.errors().is_empty(),
1167 "stale OpenSSL errors after client creation: {errors}"
1168 );
1169 }
1170
1171 #[mz_ore::test]
1172 fn test_connection_rule_pattern_matches() {
1173 let p = ConnectionRulePattern {
1174 prefix_wildcard: false,
1175 literal_match: "broker:9092".to_string(),
1176 suffix_wildcard: false,
1177 };
1178 assert!(p.matches("broker:9092"));
1179 assert!(!p.matches("other:9092"));
1180
1181 let p = ConnectionRulePattern {
1182 prefix_wildcard: true,
1183 literal_match: ":9092".to_string(),
1184 suffix_wildcard: false,
1185 };
1186 assert!(p.matches("any-host:9092"));
1187 assert!(!p.matches("broker:9093"));
1188
1189 let p = ConnectionRulePattern {
1190 prefix_wildcard: false,
1191 literal_match: "broker:".to_string(),
1192 suffix_wildcard: true,
1193 };
1194 assert!(p.matches("broker:9092"));
1195 assert!(!p.matches("other:9092"));
1196
1197 let p = ConnectionRulePattern {
1198 prefix_wildcard: true,
1199 literal_match: "broker".to_string(),
1200 suffix_wildcard: true,
1201 };
1202 assert!(p.matches("my-broker-host:1234"));
1203 assert!(!p.matches("other:9092"));
1204 }
1205}