Skip to main content

mz_kafka_util/
client.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10//! Helpers for working with Kafka's client API.
11
12use 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
48/// A reasonable default timeout when refreshing topic metadata. This is configured
49/// at a source level.
50// 30s may seem infrequent, but the default is 5m. More frequent metadata
51// refresh rates are surprising to Kafka users, as topic partition counts hardly
52// ever change in production.
53pub const DEFAULT_TOPIC_METADATA_REFRESH_INTERVAL: Duration = Duration::from_secs(30);
54
55/// A `ClientContext` implementation that uses `tracing` instead of `log`
56/// macros.
57///
58/// All code in Materialize that constructs Kafka clients should use this
59/// context or a custom context that delegates the `log` and `error` methods to
60/// this implementation.
61pub struct MzClientContext {
62    /// The last observed error log, if any.
63    error_tx: Sender<MzKafkaError>,
64    /// A tokio watch that retains the last statistics received by rdkafka and provides async
65    /// notifications to anyone interested in subscribing.
66    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    /// Constructs a new client context and returns an mpsc `Receiver` that can be used to learn
77    /// about librdkafka errors.
78    // `crossbeam` channel receivers can be cloned, but this is intended to be used as a mpsc,
79    // until we upgrade to `1.72` and the std mpsc sender is `Sync`.
80    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    /// Creates a tokio Watch subscription for statistics reported by librdkafka. It is necessary
91    /// that the `statistics.ms.interval` is set for this stream to contain any values.
92    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        // If no one cares about errors we drop them on the floor
105        let _ = self.error_tx.send(err);
106    }
107}
108
109/// A structured error type for errors reported by librdkafka through its logs.
110#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
111pub enum MzKafkaError {
112    /// Invalid username or password
113    #[error("Invalid username or password")]
114    InvalidCredentials,
115    /// Missing CA certificate
116    #[error("Invalid CA certificate")]
117    InvalidCACertificate,
118    /// Broker might require SSL encryption
119    #[error("Disconnected during handshake; broker might require SSL encryption")]
120    SSLEncryptionMaybeRequired,
121    /// Broker does not support SSL connections
122    #[error("Broker does not support SSL connections")]
123    SSLUnsupported,
124    /// Broker did not provide a certificate
125    #[error("Broker did not provide a certificate")]
126    BrokerCertificateMissing,
127    /// Failed to verify broker certificate
128    #[error("Failed to verify broker certificate")]
129    InvalidBrokerCertificate,
130    /// Connection reset
131    #[error("Connection reset: {0}")]
132    ConnectionReset(String),
133    /// Connection timeout
134    #[error("Connection timeout")]
135    ConnectionTimeout,
136    /// Failed to resolve hostname
137    #[error("Failed to resolve hostname")]
138    HostnameResolutionFailed,
139    /// Unsupported SASL mechanism
140    #[error("Unsupported SASL mechanism")]
141    UnsupportedSASLMechanism,
142    /// Unsupported broker version
143    #[error("Unsupported broker version")]
144    UnsupportedBrokerVersion,
145    /// Connection to broker failed
146    #[error("Broker transport failure")]
147    BrokerTransportFailure,
148    /// All brokers down
149    #[error("All brokers down")]
150    AllBrokersDown,
151    /// SASL authentication required
152    #[error("SASL authentication required")]
153    SaslAuthenticationRequired,
154    /// SASL authentication required
155    #[error("SASL authentication failed")]
156    SaslAuthenticationFailed,
157    /// SSL authentication required
158    #[error("SSL authentication required")]
159    SslAuthenticationRequired,
160    /// Unknown topic or partition
161    #[error("Unknown topic or partition")]
162    UnknownTopicOrPartition,
163    /// An internal kafka error
164    #[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        // Sniff out log messages that indicate errors.
233        //
234        // We consider any event at error, critical, alert, or emergency level,
235        // for self explanatory reasons. We also consider any event with a
236        // facility of `FAIL`. librdkafka often uses info or warn level for
237        // these `FAIL` events, but as they always indicate a failure to connect
238        // to a broker we want to always treat them as errors.
239        if matches!(level, Emerg | Alert | Critical | Error) || fac == "FAIL" {
240            self.record_error(log_message);
241        }
242
243        // Copied from https://docs.rs/rdkafka/0.28.0/src/rdkafka/client.rs.html#58-79
244        // but using `tracing`
245        match level {
246            Emerg | Alert | Critical | Error => {
247                // We downgrade error messages to `warn!` level to avoid
248                // sending the errors to Sentry. Most errors are customer
249                // configuration problems that are not appropriate to send to
250                // Sentry.
251                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        // Refer to the comment in the `log` callback.
267        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/// The address of a Kafka broker.
285#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd)]
286pub struct BrokerAddr {
287    /// The broker's hostname.
288    pub host: String,
289    /// The broker's port.
290    pub port: u16,
291}
292
293impl BrokerAddr {
294    /// Attempt to resolve this broker address into a list of socket addresses.
295    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/// Rewrites a broker address.
301///
302/// For use with [`TunnelingClientContext`].
303#[derive(Debug, Clone)]
304pub struct BrokerRewrite {
305    /// The rewritten hostname.
306    pub host: String,
307    /// The rewritten port.
308    ///
309    /// If unspecified, the broker's original port is left unchanged.
310    pub port: Option<u16>,
311}
312
313impl BrokerRewrite {
314    /// Apply the rewrite to this broker address.
315    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        // This ensures the ssh tunnel is not shutdown.
328        ManagedSshTunnelHandle,
329    ),
330    /// For _default_ ssh tunnels, we store an error if _creation_
331    /// of the tunnel failed, so that `tunnel_status` can return it.
332    FailedDefaultSshTunnel(String),
333}
334
335#[derive(Clone)]
336/// Parsed from a string, with optional leading and trailing '*' wildcards.
337pub struct ConnectionRulePattern {
338    /// If true, allow any combination of characters before the literal match.
339    pub prefix_wildcard: bool,
340    /// We expect the broker's host:port to match these characters in their entirety.
341    pub literal_match: String,
342    /// If true, allow any combination of characters after the literal match.
343    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    /// Does this "{host}:{port}" address fit the pattern?
361    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)]
377/// Given a host address, map it to a different host.
378pub struct HostMappingRules {
379    /// Map matching hosts to a different host. First applicable rule wins.
380    pub rules: Vec<(ConnectionRulePattern, BrokerRewrite)>,
381}
382
383impl HostMappingRules {
384    /// Rewrite this broker address according to the rules. Returns `None` when
385    /// no rule matches.
386    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/// Tunneling clients
408/// used for re-writing ports / hosts
409#[derive(Clone)]
410pub enum TunnelConfig {
411    /// Tunnel config option for SSH tunnels
412    Ssh(SshTunnelConfig),
413    /// Re-writes internal hosts using the value, used for privatelink
414    StaticHost(String),
415    /// Re-writes internal hosts according to an ordered list of rules, also used for privatelink
416    Rules(HostMappingRules),
417    /// Performs no re-writes
418    None,
419}
420
421/// A client context that supports rewriting broker addresses.
422#[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    /// Constructs a new context that wraps `inner`.
436    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    /// Adds the default broker rewrite rule.
457    ///
458    /// Connections to brokers that aren't specified in other rewrites will be rewritten to connect to
459    /// `rewrite_host` and `rewrite_port` instead.
460    pub fn set_default_tunnel(&mut self, tunnel: TunnelConfig) {
461        self.default_tunnel = tunnel;
462    }
463
464    /// Adds an SSH tunnel for a specific broker.
465    ///
466    /// Overrides the existing SSH tunnel or rewrite for this broker, if any.
467    ///
468    /// This tunnel allows the rewrite to evolve over time, for example, if
469    /// the ssh tunnel's address changes if it fails and restarts.
470    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    /// Adds a broker rewrite rule.
493    ///
494    /// Overrides the existing SSH tunnel or rewrite for this broker, if any.
495    ///
496    /// `rewrite` is `BrokerRewrite` that specifies how to rewrite the address for `broker`.
497    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    /// Returns a reference to the wrapped context.
503    pub fn inner(&self) -> &C {
504        &self.inner
505    }
506
507    /// Returns a _consolidated_ `SshTunnelStatus` that communicates the status
508    /// of all active ssh tunnels `self` knows about.
509    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        // NOTE(benesch): We abuse the `TunnelingClientContext` to handle AWS
549        // IAM authentication because it's used in exactly the right places and
550        // already has a handle to the Tokio runtime. It might be slightly
551        // cleaner to have a separate `AwsIamAuthenticatingClientContext`, but
552        // that would be quite a bit of additional plumbing.
553
554        // NOTE(benesch): at the moment, the only OAUTHBEARER authentication we
555        // support is AWS IAM, so we can assume that if this method is invoked
556        // AWS IAM is desired. We may need to generalize this in the future.
557
558        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    /// Look up the broker's address in our book of rewrites.
586    /// If we've already rewritten it before, reuse the existing rewrite.
587    /// Otherwise, use our "default tunnel" rewriting strategy to attempt to rewrite this broker's address
588    /// and record it in the book of rewrites.
589    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                    // The port for this can change over time, as the ssh tunnel is maintained through
596                    // errors.
597                    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            // No (successful) broker address rewrite exists yet.
627            None | Some(BrokerRewriteHandle::FailedDefaultSshTunnel(_)) => {
628                // "Default tunnel" is actually the configured rewriting strategy used for brokers we haven't already rewritten.
629                match &self.default_tunnel {
630                    // This "default tunnel" is actually a default tunnel.
631                    // Try connecting so we have a valid rewrite for thsi broker address.
632                    TunnelConfig::Ssh(default_tunnel) => {
633                        // Multiple users could all run `connect` at the same time; only one ssh
634                        // tunnel will ever be connected, and only one will be inserted into the
635                        // map.
636                        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                            // Use the tunnel we just created, but only if nobody beat us in the race.
649                            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                            // We couldn't connect. Someone else will have to try again.
672                            Err(e) => {
673                                warn!(
674                                    "failed to create ssh tunnel for {:?}: {}",
675                                    addr,
676                                    e.display_with_causes()
677                                );
678
679                                // Write an error if no one else has already written one.
680                                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                    // Our rewrite strategy is to use a specific host, e.g. a PrivateLink endpoint.
695                    TunnelConfig::StaticHost(host) => (host.as_str(), port)
696                        .to_socket_addrs()
697                        .map(|addrs| addrs.collect()),
698                    // Rewrite according to the routing rules.
699                    TunnelConfig::Rules(rules) => {
700                        // If no rules match, just use the address as-is.
701                        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                    // We leave the broker's address as it is.
720                    TunnelConfig::None => {
721                        (host, port).to_socket_addrs().map(|addrs| addrs.collect())
722                    }
723                }
724            }
725            // This broker's address was already rewritten. Reuse the existing rewrite.
726            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
798/// Id of a partition in a topic.
799pub type PartitionId = i32;
800
801/// The error returned by [`get_partitions`].
802#[derive(Debug, thiserror::Error)]
803pub enum GetPartitionsError {
804    /// The specified topic does not exist.
805    #[error("Topic does not exist")]
806    TopicDoesNotExist,
807    /// A Kafka error.
808    #[error(transparent)]
809    Kafka(#[from] KafkaError),
810    /// An unstructured error.
811    #[error(transparent)]
812    Other(#[from] anyhow::Error),
813}
814
815/// Retrieve number of partitions for a given `topic` using the given `client`
816pub 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
865/// Default to true as they have no downsides <https://github.com/confluentinc/librdkafka/issues/283>.
866pub const DEFAULT_KEEPALIVE: bool = true;
867/// The `rdkafka` default.
868/// - <https://github.com/confluentinc/librdkafka/blob/master/CONFIGURATION.md>
869pub const DEFAULT_SOCKET_TIMEOUT: Duration = Duration::from_secs(60);
870/// Increased from the rdkafka default
871/// - <https://github.com/confluentinc/librdkafka/blob/master/CONFIGURATION.md>
872pub const DEFAULT_TRANSACTION_TIMEOUT: Duration = Duration::from_secs(600);
873/// The `rdkafka` default.
874/// - <https://github.com/confluentinc/librdkafka/blob/master/CONFIGURATION.md>
875pub const DEFAULT_SOCKET_CONNECTION_SETUP_TIMEOUT: Duration = Duration::from_secs(30);
876/// A reasonable default timeout when fetching metadata or partitions.
877pub const DEFAULT_FETCH_METADATA_TIMEOUT: Duration = Duration::from_secs(10);
878/// The timeout for reading records from the progress topic. Set to something slightly longer than
879/// the idle transaction timeout (60s) to wait out any stuck producers.
880pub const DEFAULT_PROGRESS_RECORD_FETCH_TIMEOUT: Duration = Duration::from_secs(90);
881
882/// Configurable timeouts for Kafka connections.
883#[derive(Copy, Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
884pub struct TimeoutConfig {
885    /// Whether or not to enable
886    pub keepalive: bool,
887    /// The timeout for network requests. Can't be more than 100ms longer than
888    /// `transaction_timeout.
889    pub socket_timeout: Duration,
890    /// The timeout for transactions.
891    pub transaction_timeout: Duration,
892    /// The timeout for setting up network connections.
893    pub socket_connection_setup_timeout: Duration,
894    /// The timeout for fetching metadata from upstream.
895    pub fetch_metadata_timeout: Duration,
896    /// The timeout for reading records from the progress topic.
897    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    /// Build a `TcpTimeoutConfig` from the given parameters. Parameters outside the supported
915    /// range are defaulted and cause an error log.
916    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        // Constrain values based on ranges here:
925        // <https://github.com/confluentinc/librdkafka/blob/master/CONFIGURATION.md>
926        //
927        // Note we error log but do not fail as this is called in a non-fallible
928        // LD-sync in the adapter.
929
930        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        // The documented max here is `300000`, but rdkafka bans `socket.timeout.ms` being more
963        // than `transaction.timeout.ms` + 100ms.
964        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
1020/// A simpler version of [`create_new_client_config`] that defaults
1021/// the `log_level` to `INFO` and should only be used in tests.
1022pub fn create_new_client_config_simple() -> ClientConfig {
1023    create_new_client_config(tracing::Level::INFO, Default::default())
1024}
1025
1026/// Build a new [`rdkafka`] [`ClientConfig`] with its `log_level` set correctly
1027/// based on the passed through [`tracing::Level`]. This level should be
1028/// determined for `target: "librdkafka"`.
1029pub 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    // WARNING WARNING WARNING
1046    //
1047    // For whatever reason, if you change this `target` to something else, this
1048    // log line might break. I (guswynn) did some extensive investigation with
1049    // the tracing folks, and we learned that this edge case only happens with
1050    // 1. a different target
1051    // 2. only this file (so far as I can tell)
1052    // 3. only in certain subscriber combinations
1053    // 4. only if the `tracing-log` feature is on.
1054    //
1055    // Our conclusion was that one of our dependencies is doing something
1056    // problematic with `log`.
1057    //
1058    // For now, this works, and prints a nice log line exactly when we want it.
1059    //
1060    // TODO(guswynn): when we can remove `tracing-log`, remove this warning
1061    tracing::debug!(target: "librdkafka", level = ?level, "Determined log level for librdkafka");
1062    config.set_log_level(level);
1063
1064    // Patch the librdkafka debug log system into the Rust `log` ecosystem. This
1065    // is a very simple integration at the moment; enabling `debug`-level logs
1066    // for the `librdkafka` target enables the full firehouse of librdkafka
1067    // debug logs. We may want to investigate finer-grained control.
1068    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
1092/// Creates a client from `config` using [`ClientConfig::create`].
1093///
1094/// Leaves the calling thread's OpenSSL error queue empty, see
1095/// [`create_with_context`].
1096pub 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
1103/// Creates a client from `config` using [`ClientConfig::create_with_context`].
1104///
1105/// Leaves the calling thread's OpenSSL error queue empty.
1106pub 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    // NOTE: librdkafka loads `ssl.ca.pem` by reading certificates until
1114    // `PEM_read_bio_X509` fails, and leaves that final `PEM routines:get_name:no
1115    // start line` error on the calling thread's OpenSSL error queue. OpenSSL's
1116    // `SSL_get_error` reports any queued error as `SSL_ERROR_SSL`, so a later
1117    // TLS read that would merely block on this thread fails with the stale
1118    // error instead. On a tokio worker this breaks unrelated connections, for
1119    // example HTTPS requests to a schema registry. librdkafka v2.15.1 clears
1120    // the queue itself (confluentinc/librdkafka#5561).
1121    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)] // unsupported operation: can't call foreign function
1156    fn test_create_with_context_clears_openssl_error_queue() {
1157        let mut config = create_new_client_config_simple();
1158        // Client creation does not connect, so the broker need not exist.
1159        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}