Skip to main content

mz_balancerd/
lib.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//! The balancerd service is a horizontally scalable, stateless, multi-tenant ingress router for
11//! pgwire and HTTPS connections.
12//!
13//! It listens on pgwire and HTTPS ports. When a new pgwire connection starts, the requested user is
14//! authenticated with frontegg from which a tenant id is returned. From that a target internal
15//! hostname is resolved to an IP address, and the connection is proxied to that address which has a
16//! running environmentd's pgwire port. When a new HTTPS connection starts, its SNI hostname is used
17//! to generate an internal hostname that is resolved to an IP address, which is similarly proxied.
18
19mod codec;
20mod dyncfgs;
21
22use std::collections::BTreeMap;
23use std::net::{IpAddr, SocketAddr};
24use std::path::PathBuf;
25use std::pin::Pin;
26use std::sync::Arc;
27use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
28use std::time::{Duration, Instant};
29
30use anyhow::Context;
31use axum::response::IntoResponse;
32use axum::{Router, routing};
33use bytes::BytesMut;
34use futures::TryFutureExt;
35use futures::stream::BoxStream;
36use hickory_resolver::config::LookupIpStrategy;
37use hickory_resolver::lookup_ip::LookupIp;
38use hickory_resolver::net::runtime::TokioRuntimeProvider;
39use hickory_resolver::proto::rr::{RData, RecordType};
40use hickory_resolver::system_conf::read_system_conf;
41use hickory_resolver::{Resolver, TokioResolver};
42use hyper::StatusCode;
43use hyper_util::rt::TokioIo;
44use launchdarkly_server_sdk as ld;
45use mz_build_info::{BuildInfo, build_info};
46use mz_dyncfg::ConfigSet;
47use mz_frontegg_auth::Authenticator as FronteggAuthentication;
48use mz_ore::cast::CastFrom;
49use mz_ore::id_gen::conn_id_org_uuid;
50use mz_ore::metrics::{ComputedGauge, ComputedUIntGauge, IntCounter, IntGauge, MetricsRegistry};
51use mz_ore::netio::AsyncReady;
52use mz_ore::now::{NowFn, SYSTEM_TIME, epoch_to_uuid_v7};
53use mz_ore::task::{JoinSetExt, spawn};
54use mz_ore::tracing::TracingHandle;
55use mz_ore::{metric, netio};
56use mz_pgwire_common::{
57    ACCEPT_SSL_ENCRYPTION, CONN_UUID_KEY, Conn, ErrorResponse, FrontendMessage,
58    FrontendStartupMessage, MAX_STARTUP_FRAME_SIZE, MZ_FORWARDED_FOR_KEY, REJECT_ENCRYPTION,
59    VERSION_3, decode_startup,
60};
61use mz_server_core::{
62    Connection, ConnectionStream, ListenerHandle, ReloadTrigger, ReloadingSslContext,
63    ReloadingTlsConfig, ServeConfig, ServeDyncfg, TlsCertConfig, TlsMode, listen,
64};
65use openssl::error::ErrorStack;
66use openssl::ssl::{NameType, Ssl, SslConnector, SslMethod, SslVerifyMode};
67use prometheus::{IntCounterVec, IntGaugeVec};
68use proxy_header::{ProxiedAddress, ProxyHeader};
69use semver::Version;
70use tokio::io::{self, AsyncRead, AsyncWrite, AsyncWriteExt};
71use tokio::net::TcpStream;
72use tokio::sync::oneshot;
73use tokio::task::JoinSet;
74use tokio_metrics::TaskMetrics;
75use tokio_openssl::SslStream;
76use tokio_postgres::error::SqlState;
77use tower::Service;
78use tracing::{debug, error, info, warn};
79use uuid::Uuid;
80
81use crate::codec::{BackendMessage, FramedConn};
82use crate::dyncfgs::{
83    INJECT_PROXY_PROTOCOL_HEADER_HTTP, MAX_CONNECTIONS, PRE_RESOLVED_TIMEOUT,
84    SIGTERM_CONNECTION_WAIT, SIGTERM_LISTEN_WAIT, has_tracing_config_update, tracing_config,
85};
86
87/// Balancer build information.
88pub const BUILD_INFO: BuildInfo = build_info!();
89
90pub struct BalancerConfig {
91    /// Info about which version of the code is running.
92    build_version: Version,
93    /// Listen address for internal HTTP health and metrics server.
94    internal_http_listen_addr: SocketAddr,
95    /// Listen address for pgwire connections.
96    pgwire_listen_addr: SocketAddr,
97    /// Listen address for HTTPS connections.
98    https_listen_addr: SocketAddr,
99    /// DNS resolver for pgwire cancellation requests
100    cancellation_resolver: CancellationResolver,
101    /// DNS resolver.
102    resolver: BalancerResolver,
103    https_sni_addr_template: String,
104    tls: Option<TlsCertConfig>,
105    internal_tls: bool,
106    metrics_registry: MetricsRegistry,
107    reload_certs: BoxStream<'static, Option<oneshot::Sender<Result<(), anyhow::Error>>>>,
108    launchdarkly_sdk_key: Option<String>,
109    config_sync_file_path: Option<PathBuf>,
110    config_sync_timeout: Duration,
111    config_sync_loop_interval: Option<Duration>,
112    cloud_provider: Option<String>,
113    cloud_provider_region: Option<String>,
114    tracing_handle: TracingHandle,
115    default_configs: Vec<(String, String)>,
116}
117
118impl BalancerConfig {
119    pub fn new(
120        build_info: &BuildInfo,
121        internal_http_listen_addr: SocketAddr,
122        pgwire_listen_addr: SocketAddr,
123        https_listen_addr: SocketAddr,
124        cancellation_resolver: CancellationResolver,
125        resolver: BalancerResolver,
126        https_sni_addr_template: String,
127        tls: Option<TlsCertConfig>,
128        internal_tls: bool,
129        metrics_registry: MetricsRegistry,
130        reload_certs: ReloadTrigger,
131        launchdarkly_sdk_key: Option<String>,
132        config_sync_file: Option<PathBuf>,
133        config_sync_timeout: Duration,
134        config_sync_loop_interval: Option<Duration>,
135        cloud_provider: Option<String>,
136        cloud_provider_region: Option<String>,
137        tracing_handle: TracingHandle,
138        default_configs: Vec<(String, String)>,
139    ) -> Self {
140        Self {
141            build_version: build_info.semver_version(),
142            internal_http_listen_addr,
143            pgwire_listen_addr,
144            https_listen_addr,
145            cancellation_resolver,
146            resolver,
147            https_sni_addr_template,
148            tls,
149            internal_tls,
150            metrics_registry,
151            reload_certs,
152            launchdarkly_sdk_key,
153            config_sync_file_path: config_sync_file,
154            config_sync_timeout,
155            config_sync_loop_interval,
156            cloud_provider,
157            cloud_provider_region,
158            tracing_handle,
159            default_configs,
160        }
161    }
162}
163
164/// Prometheus monitoring metrics.
165#[derive(Debug)]
166pub struct BalancerMetrics {
167    _uptime: ComputedGauge,
168}
169
170impl BalancerMetrics {
171    /// Returns a new [BalancerMetrics] instance connected to the registry in cfg.
172    pub fn new(cfg: &BalancerConfig) -> Self {
173        let start = Instant::now();
174        let uptime = cfg.metrics_registry.register_computed_gauge(
175            metric!(
176                name: "mz_balancer_metadata_seconds",
177                help: "server uptime, labels are build metadata",
178                const_labels: {
179                    "version" => cfg.build_version,
180                    "build_type" => if cfg!(release) { "release" } else { "debug" }
181                },
182            ),
183            move || start.elapsed().as_secs_f64(),
184        );
185        BalancerMetrics { _uptime: uptime }
186    }
187}
188
189pub struct BalancerService {
190    cfg: BalancerConfig,
191    pub pgwire: (ListenerHandle, Pin<Box<dyn ConnectionStream>>),
192    pub https: (ListenerHandle, Pin<Box<dyn ConnectionStream>>),
193    pub internal_http: (ListenerHandle, Pin<Box<dyn ConnectionStream>>),
194    _metrics: BalancerMetrics,
195    configs: ConfigSet,
196}
197
198impl BalancerService {
199    pub async fn new(cfg: BalancerConfig) -> Result<Self, anyhow::Error> {
200        let pgwire = listen(&cfg.pgwire_listen_addr).await?;
201        let https = listen(&cfg.https_listen_addr).await?;
202        let internal_http = listen(&cfg.internal_http_listen_addr).await?;
203        let metrics = BalancerMetrics::new(&cfg);
204        let mut configs = ConfigSet::default();
205        configs = dyncfgs::all_dyncfgs(configs);
206        dyncfgs::set_defaults(&configs, cfg.default_configs.clone())?;
207        let tracing_handle = cfg.tracing_handle.clone();
208        // Configure dyncfg sync
209        match (
210            cfg.launchdarkly_sdk_key.as_deref(),
211            cfg.config_sync_file_path.as_deref(),
212        ) {
213            (Some(key), None) => {
214                let _ = mz_dyncfg_launchdarkly::sync_launchdarkly_to_configset(
215                    configs.clone(),
216                    &BUILD_INFO,
217                    |builder| {
218                        let region = cfg
219                            .cloud_provider_region
220                            .clone()
221                            .unwrap_or_else(|| String::from("unknown"));
222                        if let Some(provider) = cfg.cloud_provider.clone() {
223                            builder.add_context(
224                                ld::ContextBuilder::new(format!(
225                                    "{}/{}/{}",
226                                    provider, region, cfg.build_version
227                                ))
228                                .kind("balancer")
229                                .set_string("provider", provider)
230                                .set_string("region", region)
231                                .set_string("version", cfg.build_version.to_string())
232                                .build()
233                                .map_err(|e| anyhow::anyhow!(e))?,
234                            );
235                        } else {
236                            builder.add_context(
237                                ld::ContextBuilder::new(format!(
238                                    "{}/{}/{}",
239                                    "unknown", region, cfg.build_version
240                                ))
241                                .anonymous(true) // exclude this user from the dashboard
242                                .kind("balancer")
243                                .set_string("provider", "unknown")
244                                .set_string("region", region)
245                                .set_string("version", cfg.build_version.to_string())
246                                .build()
247                                .map_err(|e| anyhow::anyhow!(e))?,
248                            );
249                        }
250                        Ok(())
251                    },
252                    Some(key),
253                    cfg.config_sync_timeout,
254                    cfg.config_sync_loop_interval,
255                    move |updates, configs| {
256                        if has_tracing_config_update(updates) {
257                            match tracing_config(configs) {
258                                Ok(parameters) => parameters.apply(&tracing_handle),
259                                Err(err) => warn!("unable to update tracing: {err}"),
260                            }
261                        }
262                    },
263                )
264                .await
265                .inspect_err(|e| warn!("LaunchDarkly sync error: {e}"));
266            }
267            (None, Some(path)) => {
268                let _ = mz_dyncfg_file::sync_file_to_configset(
269                    configs.clone(),
270                    path,
271                    cfg.config_sync_timeout,
272                    cfg.config_sync_loop_interval,
273                    move |updates, configs| {
274                        if has_tracing_config_update(updates) {
275                            match tracing_config(configs) {
276                                Ok(parameters) => parameters.apply(&tracing_handle),
277                                Err(err) => warn!("unable to update tracing: {err}"),
278                            }
279                        }
280                    },
281                )
282                .await
283                // If there's an Error, log but continue anyway. If LD is down
284                // we have no way of fetching the previous value of the flag
285                // (unlike the adapter, but it has a durable catalog). The
286                // ConfigSet defaults have been chosen to be good enough if this
287                // is the case.
288                .inspect_err(|e| warn!("File config sync error: {e}"));
289            }
290            (Some(_), Some(_)) => panic!(
291                "must provide either config_sync_file_path or launchdarkly_sdk_key for config syncing",
292            ),
293            (None, None) => {}
294        };
295        Ok(Self {
296            cfg,
297            pgwire,
298            https,
299            internal_http,
300            _metrics: metrics,
301            configs,
302        })
303    }
304
305    pub async fn serve(self) -> Result<(), anyhow::Error> {
306        let (pgwire_tls, https_tls) = match &self.cfg.tls {
307            Some(tls) => {
308                let context = tls.reloading_context(self.cfg.reload_certs)?;
309                (
310                    Some(ReloadingTlsConfig {
311                        context: context.clone(),
312                        mode: TlsMode::Require,
313                    }),
314                    Some(context),
315                )
316            }
317            None => (None, None),
318        };
319
320        // Built once and shared by every upstream connection, since building a connector loads
321        // and parses the system CA bundle.
322        let internal_tls = self
323            .cfg
324            .internal_tls
325            .then(internal_tls_connector)
326            .transpose()
327            .context("building internal TLS connector")?;
328
329        let metrics = ServerMetricsConfig::register_into(&self.cfg.metrics_registry);
330        let limiter = ConnectionLimiter::new(&self.cfg.metrics_registry, self.configs.clone());
331
332        let mut set = JoinSet::new();
333        let mut server_handles = Vec::new();
334        let pgwire_addr = self.pgwire.0.local_addr();
335        let https_addr = self.https.0.local_addr();
336        let internal_http_addr = self.internal_http.0.local_addr();
337
338        // The HTTPS balancer always resolves through a TenantDnsResolver. In
339        // multi-tenant mode it shares the pgwire resolver so both listeners use
340        // one resolver. In static mode pgwire does not resolve through it, so
341        // HTTPS gets its own.
342        let shared_dns = match self.cfg.resolver.shared_dns() {
343            Some(dns) => dns,
344            None => Arc::new(TenantDnsResolver::new()?),
345        };
346
347        {
348            let pgwire = PgwireBalancer {
349                resolver: Arc::new(self.cfg.resolver),
350                cancellation_resolver: Arc::new(self.cfg.cancellation_resolver),
351                tls: pgwire_tls,
352                internal_tls: internal_tls.clone(),
353                metrics: ServerMetrics::new(metrics.clone(), "pgwire"),
354                limiter: Arc::clone(&limiter),
355                configs: self.configs.clone(),
356                now: SYSTEM_TIME.clone(),
357            };
358            let (handle, stream) = self.pgwire;
359            server_handles.push(handle);
360            set.spawn_named(|| "pgwire_stream", {
361                let config_set = self.configs.clone();
362                async move {
363                    mz_server_core::serve(ServeConfig {
364                        server: pgwire,
365                        conns: stream,
366                        dyncfg: Some(ServeDyncfg {
367                            config_set,
368                            sigterm_wait_config: &SIGTERM_CONNECTION_WAIT,
369                        }),
370                    })
371                    .await;
372                    warn!("pgwire server exited");
373                }
374            });
375        }
376        {
377            let Some((addr, port)) = self.cfg.https_sni_addr_template.split_once(':') else {
378                panic!("expected port in https_addr_template");
379            };
380            let port: u16 = port.parse().expect("unexpected port");
381
382            let https = HttpsBalancer {
383                resolver: shared_dns,
384                tls: https_tls,
385                resolve_template: Arc::from(addr),
386                port,
387                metrics: Arc::from(ServerMetrics::new(metrics, "https")),
388                limiter,
389                configs: self.configs.clone(),
390                internal_tls,
391            };
392            let (handle, stream) = self.https;
393            server_handles.push(handle);
394            set.spawn_named(|| "https_stream", {
395                let config_set = self.configs.clone();
396                async move {
397                    mz_server_core::serve(ServeConfig {
398                        server: https,
399                        conns: stream,
400                        dyncfg: Some(ServeDyncfg {
401                            config_set,
402                            sigterm_wait_config: &SIGTERM_CONNECTION_WAIT,
403                        }),
404                    })
405                    .await;
406                    warn!("https server exited");
407                }
408            });
409        }
410        {
411            let router = Router::new()
412                .route(
413                    "/metrics",
414                    routing::get(move |headers: axum::http::HeaderMap| async move {
415                        mz_http_util::handle_prometheus(&self.cfg.metrics_registry, headers).await
416                    }),
417                )
418                .route(
419                    "/api/livez",
420                    routing::get(mz_http_util::handle_liveness_check),
421                )
422                .route("/api/readyz", routing::get(handle_readiness_check));
423            let internal_http = InternalHttpServer { router };
424            let (handle, stream) = self.internal_http;
425            server_handles.push(handle);
426            set.spawn_named(|| "internal_http_stream", async move {
427                mz_server_core::serve(ServeConfig {
428                    server: internal_http,
429                    conns: stream,
430                    // Disable graceful termination because our internal
431                    // monitoring keeps persistent HTTP connections open.
432                    dyncfg: None,
433                })
434                .await;
435                warn!("internal_http server exited");
436            });
437        }
438        #[cfg(unix)]
439        {
440            let mut sigterm =
441                tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())?;
442            set.spawn_named(|| "sigterm_handler", async move {
443                sigterm.recv().await;
444                let wait = SIGTERM_LISTEN_WAIT.get(&self.configs);
445                warn!("received signal TERM - delaying for {:?}!", wait);
446                tokio::time::sleep(wait).await;
447                warn!("sigterm delay complete, dropping server handles");
448                drop(server_handles);
449            });
450        }
451
452        println!("balancerd {} listening...", BUILD_INFO.human_version(None));
453        println!(" TLS enabled: {}", self.cfg.tls.is_some());
454        println!(" pgwire address: {}", pgwire_addr);
455        println!(" HTTPS address: {}", https_addr);
456        println!(" internal HTTP address: {}", internal_http_addr);
457
458        // Wait for all tasks to exit, which can happen on SIGTERM.
459        while let Some(res) = set.join_next().await {
460            if let Err(err) = res {
461                error!("serving task failed: {err}")
462            }
463        }
464        Ok(())
465    }
466}
467
468#[allow(clippy::unused_async)]
469async fn handle_readiness_check() -> impl IntoResponse {
470    (StatusCode::OK, "ready")
471}
472
473struct InternalHttpServer {
474    router: Router,
475}
476
477impl mz_server_core::Server for InternalHttpServer {
478    const NAME: &'static str = "internal_http";
479
480    // TODO(jkosh44) consider forwarding the connection UUID to the adapter.
481    fn handle_connection(
482        &self,
483        conn: Connection,
484        _tokio_metrics_intervals: impl Iterator<Item = TaskMetrics> + Send + 'static,
485    ) -> mz_server_core::ConnectionHandler {
486        let router = self.router.clone();
487        let service = hyper::service::service_fn(move |req| router.clone().call(req));
488        let conn = TokioIo::new(conn);
489
490        Box::pin(async {
491            let http = hyper::server::conn::http1::Builder::new();
492            http.serve_connection(conn, service).err_into().await
493        })
494    }
495}
496
497/// Builds the connector for TLS to environmentd.
498fn internal_tls_connector() -> Result<SslConnector, ErrorStack> {
499    let mut builder = SslConnector::builder(SslMethod::tls())?;
500    // environmentd doesn't yet have a cert we trust, so for now disable verification.
501    builder.set_verify(SslVerifyMode::NONE);
502    Ok(builder.build())
503}
504
505/// Wraps an IntGauge and automatically `inc`s on init and `drop`s on drop. Callers should not call
506/// `inc().`. Useful for handling multiple task exit points, for example in the case of a panic.
507struct GaugeGuard {
508    gauge: IntGauge,
509}
510
511impl From<IntGauge> for GaugeGuard {
512    fn from(gauge: IntGauge) -> Self {
513        let _self = Self { gauge };
514        _self.gauge.inc();
515        _self
516    }
517}
518
519impl Drop for GaugeGuard {
520    fn drop(&mut self) {
521        self.gauge.dec();
522    }
523}
524
525#[derive(Clone, Debug)]
526struct ServerMetricsConfig {
527    connection_status: IntCounterVec,
528    active_connections: IntGaugeVec,
529    pre_resolved_connections: IntGaugeVec,
530    pre_resolved_timeouts: IntCounterVec,
531    tenant_connections: IntGaugeVec,
532    tenant_connection_rx: IntCounterVec,
533    tenant_connection_tx: IntCounterVec,
534    tenant_pgwire_sni_count: IntCounterVec,
535}
536
537impl ServerMetricsConfig {
538    fn register_into(registry: &MetricsRegistry) -> Self {
539        let connection_status = registry.register(metric!(
540            name: "mz_balancer_connection_status",
541            help: "Count of completed network connections, by status",
542            var_labels: ["source", "status"],
543        ));
544        let active_connections = registry.register(metric!(
545            name: "mz_balancer_connection_active",
546            help: "Count of currently open network connections.",
547            var_labels: ["source"],
548        ));
549        let pre_resolved_connections = registry.register(metric!(
550            name: "mz_balancer_pre_resolved_connection_active",
551            help: "Count of open network connections that have not yet resolved a backend.",
552            var_labels: ["source"],
553        ));
554        let pre_resolved_timeouts = registry.register(metric!(
555            name: "mz_balancer_pre_resolved_timeout_total",
556            help: "Count of connections closed for not resolving a backend within the pre-resolved timeout.",
557            var_labels: ["source"],
558        ));
559        let tenant_connections = registry.register(metric!(
560            name: "mz_balancer_tenant_connection_active",
561            help: "Count of opened network connections by tenant.",
562            var_labels: ["source",  "tenant"]
563        ));
564        let tenant_connection_rx = registry.register(metric!(
565            name: "mz_balancer_tenant_connection_rx",
566            help: "Number of bytes received from a client for a tenant.",
567            var_labels: ["source", "tenant"],
568        ));
569        let tenant_connection_tx = registry.register(metric!(
570            name: "mz_balancer_tenant_connection_tx",
571            help: "Number of bytes sent to a client for a tenant.",
572            var_labels: ["source", "tenant"],
573        ));
574        let tenant_pgwire_sni_count = registry.register(metric!(
575            name: "mz_balancer_tenant_pgwire_sni_count",
576            help: "Count of pgwire connections that have and do not have SNI available per tenant.",
577            var_labels: ["tenant", "has_sni"],
578        ));
579        Self {
580            connection_status,
581            active_connections,
582            pre_resolved_connections,
583            pre_resolved_timeouts,
584            tenant_connections,
585            tenant_connection_rx,
586            tenant_connection_tx,
587            tenant_pgwire_sni_count,
588        }
589    }
590}
591
592#[derive(Clone, Debug)]
593struct ServerMetrics {
594    inner: ServerMetricsConfig,
595    source: &'static str,
596}
597
598impl ServerMetrics {
599    fn new(inner: ServerMetricsConfig, source: &'static str) -> Self {
600        let self_ = Self { inner, source };
601
602        // Pre-initialize labels we are planning to use to ensure they are all always emitted as
603        // time series.
604        self_.connection_status(false);
605        self_.connection_status(true);
606        drop(self_.active_connections());
607
608        self_
609    }
610
611    fn connection_status(&self, is_ok: bool) -> IntCounter {
612        self.inner
613            .connection_status
614            .with_label_values(&[self.source, Self::status_label(is_ok)])
615    }
616
617    fn active_connections(&self) -> GaugeGuard {
618        self.inner
619            .active_connections
620            .with_label_values(&[self.source])
621            .into()
622    }
623
624    fn pre_resolved_connections(&self) -> GaugeGuard {
625        self.inner
626            .pre_resolved_connections
627            .with_label_values(&[self.source])
628            .into()
629    }
630
631    fn pre_resolved_timeout(&self) -> IntCounter {
632        self.inner
633            .pre_resolved_timeouts
634            .with_label_values(&[self.source])
635    }
636
637    fn tenant_connections(&self, tenant: &str) -> GaugeGuard {
638        self.inner
639            .tenant_connections
640            .with_label_values(&[self.source, tenant])
641            .into()
642    }
643
644    fn tenant_connections_rx(&self, tenant: &str) -> IntCounter {
645        self.inner
646            .tenant_connection_rx
647            .with_label_values(&[self.source, tenant])
648    }
649
650    fn tenant_connections_tx(&self, tenant: &str) -> IntCounter {
651        self.inner
652            .tenant_connection_tx
653            .with_label_values(&[self.source, tenant])
654    }
655
656    fn tenant_pgwire_sni_count(&self, tenant: &str, has_sni: bool) -> IntCounter {
657        self.inner
658            .tenant_pgwire_sni_count
659            .with_label_values(&[tenant, &has_sni.to_string()])
660    }
661
662    fn status_label(is_ok: bool) -> &'static str {
663        if is_ok { "success" } else { "error" }
664    }
665}
666
667/// Ceiling on the number of client connections proxied at once, shared by the pgwire and HTTPS
668/// listeners.
669///
670/// balancerd's memory use scales with the number of connections it proxies. Without a ceiling it
671/// keeps accepting until the container is OOM killed, which drops every established connection.
672/// Refusing new connections instead sheds only the load we cannot serve.
673#[derive(Debug)]
674struct ConnectionLimiter {
675    configs: ConfigSet,
676    active: AtomicU32,
677    /// Whether connections are currently being refused, so that reaching and clearing the limit
678    /// are logged once each rather than once per connection.
679    limited: AtomicBool,
680    rejected: IntCounter,
681    _limit: ComputedUIntGauge,
682}
683
684impl ConnectionLimiter {
685    fn new(registry: &MetricsRegistry, configs: ConfigSet) -> Arc<Self> {
686        let rejected = registry.register(metric!(
687            name: "mz_balancer_connection_rejected_total",
688            help: "Count of connections refused because the connection limit was reached.",
689        ));
690        let limit = registry.register_computed_gauge(
691            metric!(
692                name: "mz_balancer_connection_limit",
693                help: "Maximum number of connections proxied at once, 0 if unlimited.",
694            ),
695            {
696                let configs = configs.clone();
697                move || u64::from(MAX_CONNECTIONS.get(&configs))
698            },
699        );
700        Arc::new(ConnectionLimiter {
701            configs,
702            active: AtomicU32::new(0),
703            limited: AtomicBool::new(false),
704            rejected,
705            _limit: limit,
706        })
707    }
708
709    /// Reserves capacity for one connection, or returns `None` if the limit has been reached, in
710    /// which case the caller must refuse the connection.
711    fn acquire(self: &Arc<Self>) -> Option<ConnectionGuard> {
712        let limit = MAX_CONNECTIONS.get(&self.configs);
713        let unlimited = limit == 0;
714        match self
715            .active
716            .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |active| {
717                (unlimited || active < limit).then(|| active + 1)
718            }) {
719            Ok(prev) => {
720                // Only clear the limited state once comfortably below the limit, so that
721                // connections churning at the ceiling do not flap the log lines.
722                if unlimited || prev + 1 <= limit - limit / 10 {
723                    if self.limited.swap(false, Ordering::Relaxed) {
724                        info!("accepting new connections again, limit is {limit}");
725                    }
726                }
727                Some(ConnectionGuard(Arc::clone(self)))
728            }
729            Err(active) => {
730                self.rejected.inc();
731                if !self.limited.swap(true, Ordering::Relaxed) {
732                    warn!(
733                        "refusing new connections: at the limit of {limit} connections \
734                        ({active} active)"
735                    );
736                }
737                None
738            }
739        }
740    }
741}
742
743/// Releases the connection reserved by [`ConnectionLimiter::acquire`] when dropped.
744struct ConnectionGuard(Arc<ConnectionLimiter>);
745
746impl Drop for ConnectionGuard {
747    fn drop(&mut self) {
748        self.0.active.fetch_sub(1, Ordering::SeqCst);
749    }
750}
751
752/// Answers the client with a final message and flushes it.
753///
754/// [`FramedConn::send`] only enqueues into a buffer, so a rejection that is not flushed reaches
755/// the client as a bare socket close rather than the error it was given. Every path that answers
756/// the client and then stops has to flush, because the connection is dropped as soon as it
757/// returns.
758async fn reject<A>(conn: &mut FramedConn<A>, err: ErrorResponse) -> Result<(), io::Error>
759where
760    A: AsyncRead + AsyncWrite + Unpin,
761{
762    conn.send(err).await?;
763    conn.flush().await
764}
765
766/// Runs the pre-resolved phase under [`PRE_RESOLVED_TIMEOUT`], counting an overrun.
767///
768/// `fut` must cover TLS negotiation, startup, authentication, and backend resolution.
769/// Proxying must run after it completes, outside the deadline. A zero timeout leaves
770/// `fut` unbounded.
771async fn under_pre_resolved_timeout<F, T, E>(
772    timeout: Duration,
773    metrics: &ServerMetrics,
774    fut: F,
775) -> Result<T, E>
776where
777    F: Future<Output = Result<T, E>>,
778    E: From<io::Error>,
779{
780    if timeout.is_zero() {
781        return fut.await;
782    }
783    match tokio::time::timeout(timeout, fut).await {
784        Ok(result) => result,
785        Err(_) => {
786            metrics.pre_resolved_timeout().inc();
787            Err(io::Error::new(
788                io::ErrorKind::TimedOut,
789                "timed out before resolving a backend",
790            )
791            .into())
792        }
793    }
794}
795
796pub enum CancellationResolver {
797    Directory(PathBuf),
798    Static(String),
799}
800
801struct PgwireBalancer {
802    /// Read for [`PRE_RESOLVED_TIMEOUT`] at connection time.
803    configs: ConfigSet,
804    tls: Option<ReloadingTlsConfig>,
805    internal_tls: Option<SslConnector>,
806    cancellation_resolver: Arc<CancellationResolver>,
807    resolver: Arc<BalancerResolver>,
808    metrics: ServerMetrics,
809    limiter: Arc<ConnectionLimiter>,
810    now: NowFn,
811}
812
813impl PgwireBalancer {
814    /// Negotiates pgwire startup and resolves a backend.
815    ///
816    /// Returns `None` after a client disconnect, a flushed rejection, or a dispatched
817    /// cancellation request.
818    async fn pre_resolve(
819        conn: Connection,
820        tls: Option<ReloadingTlsConfig>,
821        resolver: &BalancerResolver,
822        cancellation_resolver: Arc<CancellationResolver>,
823        conn_uuid: Uuid,
824        metrics: &ServerMetrics,
825    ) -> Result<
826        Option<(
827            FramedConn<Connection>,
828            ResolvedAddr,
829            BTreeMap<String, String>,
830        )>,
831        anyhow::Error,
832    > {
833        let peer_addr = conn.peer_addr();
834        let mut conn = Conn::Unencrypted(conn);
835        loop {
836            let message = decode_startup(&mut conn, MAX_STARTUP_FRAME_SIZE).await?;
837            conn = match message {
838                // Clients sometimes hang up during the startup sequence, e.g.
839                // because they receive an unacceptable response to an
840                // `SslRequest`. This is considered a graceful termination.
841                None => return Ok(None),
842
843                Some(FrontendStartupMessage::Startup {
844                    version,
845                    mut params,
846                }) => {
847                    let mut conn = FramedConn::new(conn);
848                    let rejected = SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION;
849                    let peer_addr = match peer_addr {
850                        Ok(addr) => addr.ip(),
851                        Err(e) => {
852                            error!("Invalid peer_addr {:?}", e);
853                            reject(
854                                &mut conn,
855                                ErrorResponse::fatal(rejected, "invalid peer address"),
856                            )
857                            .await?;
858                            return Ok(None);
859                        }
860                    };
861                    debug!(
862                        %conn_uuid, %peer_addr,
863                        "starting new pgwire connection in balancer",
864                    );
865                    let prev = params.insert(CONN_UUID_KEY.to_string(), conn_uuid.to_string());
866                    if prev.is_some() {
867                        reject(
868                            &mut conn,
869                            ErrorResponse::fatal(
870                                rejected,
871                                format!("invalid parameter '{CONN_UUID_KEY}'"),
872                            ),
873                        )
874                        .await?;
875                        return Ok(None);
876                    }
877
878                    let forwarded_for = params.insert(
879                        MZ_FORWARDED_FOR_KEY.to_string(),
880                        peer_addr.to_string().clone(),
881                    );
882                    if let Some(_) = forwarded_for {
883                        reject(
884                            &mut conn,
885                            ErrorResponse::fatal(
886                                rejected,
887                                format!("invalid parameter '{MZ_FORWARDED_FOR_KEY}'"),
888                            ),
889                        )
890                        .await?;
891                        return Ok(None);
892                    };
893
894                    let Some(resolved) = Self::resolve_destination(
895                        &mut conn,
896                        version,
897                        &params,
898                        resolver,
899                        tls.map(|tls| tls.mode),
900                        metrics,
901                    )
902                    .await?
903                    else {
904                        return Ok(None);
905                    };
906                    return Ok(Some((conn, resolved, params)));
907                }
908
909                Some(FrontendStartupMessage::CancelRequest {
910                    conn_id,
911                    secret_key,
912                }) => {
913                    spawn(|| "cancel request", async move {
914                        cancel_request(conn_id, secret_key, &cancellation_resolver).await;
915                    });
916                    // Do not wait on cancel requests to return because cancellation is best
917                    // effort.
918                    return Ok(None);
919                }
920
921                Some(FrontendStartupMessage::SslRequest) => match (conn, &tls) {
922                    (Conn::Unencrypted(mut conn), Some(tls)) => {
923                        conn.write_all(&[ACCEPT_SSL_ENCRYPTION]).await?;
924                        let mut ssl_stream = SslStream::new(Ssl::new(&tls.context.get())?, conn)?;
925                        if let Err(e) = Pin::new(&mut ssl_stream).accept().await {
926                            let _ = ssl_stream.get_mut().shutdown().await;
927                            return Err(e.into());
928                        }
929                        Conn::Ssl(ssl_stream)
930                    }
931                    (mut conn, _) => {
932                        conn.write_all(&[REJECT_ENCRYPTION]).await?;
933                        conn
934                    }
935                },
936
937                Some(FrontendStartupMessage::GssEncRequest) => {
938                    conn.write_all(&[REJECT_ENCRYPTION]).await?;
939                    conn
940                }
941            }
942        }
943    }
944
945    #[mz_ore::instrument(level = "debug")]
946    /// Resolves a destination for the connection, or answers the client saying why it cannot.
947    ///
948    /// `Ok(None)` means the rejection has been flushed and the connection is finished.
949    async fn resolve_destination<'a, A>(
950        conn: &'a mut FramedConn<A>,
951        version: i32,
952        params: &BTreeMap<String, String>,
953        resolver: &BalancerResolver,
954        tls_mode: Option<TlsMode>,
955        metrics: &ServerMetrics,
956    ) -> Result<Option<ResolvedAddr>, io::Error>
957    where
958        A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
959    {
960        if version != VERSION_3 {
961            reject(
962                conn,
963                ErrorResponse::fatal(
964                    SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION,
965                    "server does not support the client's requested protocol version",
966                ),
967            )
968            .await?;
969            return Ok(None);
970        }
971
972        let Some(user) = params.get("user") else {
973            reject(
974                conn,
975                ErrorResponse::fatal(
976                    SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION,
977                    "user parameter required",
978                ),
979            )
980            .await?;
981            return Ok(None);
982        };
983
984        if let Err(err) = conn.inner().ensure_tls_compatibility(&tls_mode) {
985            reject(conn, err).await?;
986            return Ok(None);
987        }
988
989        let resolved = match resolver.resolve(conn, user, metrics).await {
990            Ok(v) => v,
991            Err(err) => {
992                let sql_state = match &err {
993                    ResolveError::InvalidPassword => SqlState::INVALID_PASSWORD,
994                    ResolveError::Client(details) => {
995                        warn!("client-caused connection failure: {details:#}");
996                        SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION
997                    }
998                    ResolveError::Upstream(details) => {
999                        warn!("upstream not available: {details:#}");
1000                        SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION
1001                    }
1002                    ResolveError::Internal(details) => {
1003                        error!("resolving connection destination: {details:#}");
1004                        SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION
1005                    }
1006                };
1007                reject(conn, ErrorResponse::fatal(sql_state, err.to_string())).await?;
1008                return Ok(None);
1009            }
1010        };
1011
1012        Ok(Some(resolved))
1013    }
1014
1015    /// Proxies a resolved connection until either side closes.
1016    async fn proxy<'a, A>(
1017        conn: &'a mut FramedConn<A>,
1018        resolved: ResolvedAddr,
1019        params: BTreeMap<String, String>,
1020        internal_tls: Option<&SslConnector>,
1021        metrics: &ServerMetrics,
1022    ) -> Result<(), io::Error>
1023    where
1024        A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
1025    {
1026        let _active_guard = resolved
1027            .tenant
1028            .as_ref()
1029            .map(|tenant| metrics.tenant_connections(tenant));
1030        let mut mz_stream =
1031            match Self::init_stream(conn, resolved.addr, resolved.password, params, internal_tls)
1032                .await
1033            {
1034                Ok(stream) => stream,
1035                Err(e) => {
1036                    error!("failed to connect to upstream server: {e}");
1037                    return conn
1038                        .send(ErrorResponse::fatal(
1039                            SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION,
1040                            "upstream server not available",
1041                        ))
1042                        .await;
1043                }
1044            };
1045
1046        let mut client_counter = CountingConn::new(conn.inner_mut());
1047
1048        // Now blindly shuffle bytes back and forth until closed.
1049        // TODO: Limit total memory use.
1050        let res = tokio::io::copy_bidirectional(&mut client_counter, &mut mz_stream).await;
1051        if let Some(tenant) = &resolved.tenant {
1052            metrics
1053                .tenant_connections_tx(tenant)
1054                .inc_by(u64::cast_from(client_counter.written));
1055            metrics
1056                .tenant_connections_rx(tenant)
1057                .inc_by(u64::cast_from(client_counter.read));
1058        }
1059        res?;
1060
1061        Ok(())
1062    }
1063
1064    #[mz_ore::instrument(level = "debug")]
1065    async fn init_stream<'a, A>(
1066        conn: &'a mut FramedConn<A>,
1067        envd_addr: SocketAddr,
1068        password: Option<String>,
1069        params: BTreeMap<String, String>,
1070        internal_tls: Option<&SslConnector>,
1071    ) -> Result<Conn<TcpStream>, anyhow::Error>
1072    where
1073        A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
1074    {
1075        let mut mz_stream = TcpStream::connect(envd_addr).await?;
1076        let mut buf = BytesMut::new();
1077
1078        let mut mz_stream = if let Some(internal_tls) = internal_tls {
1079            FrontendStartupMessage::SslRequest.encode(&mut buf)?;
1080            mz_stream.write_all(&buf).await?;
1081            buf.clear();
1082            let mut maybe_ssl_request_response = [0u8; 1];
1083            let nread =
1084                netio::read_exact_or_eof(&mut mz_stream, &mut maybe_ssl_request_response).await?;
1085            if nread == 1 && maybe_ssl_request_response == [ACCEPT_SSL_ENCRYPTION] {
1086                // do a TLS handshake
1087                let mut ssl = internal_tls.configure()?.into_ssl(&envd_addr.to_string())?;
1088                ssl.set_connect_state();
1089                Conn::Ssl(SslStream::new(ssl, mz_stream)?)
1090            } else {
1091                Conn::Unencrypted(mz_stream)
1092            }
1093        } else {
1094            Conn::Unencrypted(mz_stream)
1095        };
1096
1097        // Send initial startup and password messages.
1098        let startup = FrontendStartupMessage::Startup {
1099            version: VERSION_3,
1100            params,
1101        };
1102        startup.encode(&mut buf)?;
1103        mz_stream.write_all(&buf).await?;
1104        let client_stream = conn.inner_mut();
1105
1106        // This early return is important in self managed with SASL mode.
1107        // The below code specifically looks for cleartext password requests, but in SASL mode
1108        // the server will send a different message type (SASLInitialResponse) that we should
1109        // not try to interpret or respond to.
1110        // "Why not? That code looks like it should fall back fine?" You may ask.
1111        // The below block unconditionally reads 9 bytes from the server. If we don't have
1112        // a password or the message isn't a cleartext password request, we forward those 9 bytes
1113        // to the client. Then we return the stream to the caller, who will continue shuffling bytes.
1114        // The problem is that with TLS enabled between balancerd <-> client, flushing the first 9 bytes
1115        // before copying bidirectionally will have the side effect of splitting the auth handshake into
1116        // two SSL records. Pgbouncer misbehaves in this scenario, and fails the connection.
1117        // PGbouncer shouldn't do this! It's a common footgun of protocols over TLS.
1118        // So common in fact that PGbouncer already hit and fixed this issue on the bouncer <-> client side:
1119        // once before: https://github.com/pgbouncer/pgbouncer/pull/1058.
1120        // We will work to upstream a fix, but in the meantime, this early return avoids the issue entirely.
1121        if password.is_none() {
1122            return Ok(mz_stream);
1123        }
1124
1125        // Read a single backend message, which may be a password request. Send ours if so.
1126        // Otherwise start shuffling bytes. message type (len 1, 'R') + message len (len 4, 8_i32) +
1127        // auth type (len 4, 3_i32).
1128        let mut maybe_auth_frame = [0; 1 + 4 + 4];
1129        let nread = netio::read_exact_or_eof(&mut mz_stream, &mut maybe_auth_frame).await?;
1130        // 'R' for auth message, 0008 for message length, 0003 for password cleartext variant.
1131        // See: https://www.postgresql.org/docs/current/protocol-message-formats.html#PROTOCOL-MESSAGE-FORMATS-AUTHENTICATIONCLEARTEXTPASSWORD
1132        const AUTH_PASSWORD_CLEARTEXT: [u8; 9] = [b'R', 0, 0, 0, 8, 0, 0, 0, 3];
1133        if nread == AUTH_PASSWORD_CLEARTEXT.len()
1134            && maybe_auth_frame == AUTH_PASSWORD_CLEARTEXT
1135            && password.is_some()
1136        {
1137            // If we got exactly a cleartext password request and have one, send it.
1138            let Some(password) = password else {
1139                unreachable!("verified some above");
1140            };
1141            let password = FrontendMessage::Password { password };
1142            buf.clear();
1143            password.encode(&mut buf)?;
1144            mz_stream.write_all(&buf).await?;
1145            mz_stream.flush().await?;
1146        } else {
1147            // Otherwise pass on the bytes we just got. This *might* even be a password request, but
1148            // we don't have a password. In which case it can be forwarded up to the client.
1149            client_stream.write_all(&maybe_auth_frame[0..nread]).await?;
1150        }
1151
1152        Ok(mz_stream)
1153    }
1154}
1155
1156impl mz_server_core::Server for PgwireBalancer {
1157    const NAME: &'static str = "pgwire_balancer";
1158
1159    fn handle_connection(
1160        &self,
1161        conn: Connection,
1162        _tokio_metrics_intervals: impl Iterator<Item = TaskMetrics> + Send + 'static,
1163    ) -> mz_server_core::ConnectionHandler {
1164        let tls = self.tls.clone();
1165        let internal_tls = self.internal_tls.clone();
1166        let resolver = Arc::clone(&self.resolver);
1167        let inner_metrics = self.metrics.clone();
1168        let outer_metrics = self.metrics.clone();
1169        let limiter = Arc::clone(&self.limiter);
1170        let pre_resolved_timeout = PRE_RESOLVED_TIMEOUT.get(&self.configs);
1171        let cancellation_resolver = Arc::clone(&self.cancellation_resolver);
1172        let conn_uuid = epoch_to_uuid_v7(&(self.now)());
1173        conn.uuid_handle().set(conn_uuid);
1174        Box::pin(async move {
1175            // TODO: Try to merge this with pgwire/server.rs to avoid the duplication. May not be
1176            // worth it.
1177            let active_guard = outer_metrics.active_connections();
1178            let result: Result<(), anyhow::Error> = async move {
1179                let Some(_conn_guard) = limiter.acquire() else {
1180                    return Ok(());
1181                };
1182                let pre_resolved = inner_metrics.pre_resolved_connections();
1183
1184                let destination = under_pre_resolved_timeout(
1185                    pre_resolved_timeout,
1186                    &inner_metrics,
1187                    Self::pre_resolve(
1188                        conn,
1189                        tls,
1190                        &resolver,
1191                        cancellation_resolver,
1192                        conn_uuid,
1193                        &inner_metrics,
1194                    ),
1195                )
1196                .await?;
1197
1198                let Some((mut conn, resolved, params)) = destination else {
1199                    return Ok(());
1200                };
1201                drop(pre_resolved);
1202
1203                PgwireBalancer::proxy(
1204                    &mut conn,
1205                    resolved,
1206                    params,
1207                    internal_tls.as_ref(),
1208                    &inner_metrics,
1209                )
1210                .await?;
1211                conn.flush().await?;
1212                Ok(())
1213            }
1214            .await;
1215            drop(active_guard);
1216            outer_metrics.connection_status(result.is_ok()).inc();
1217            Ok(())
1218        })
1219    }
1220}
1221
1222// A struct that counts bytes exchanged.
1223struct CountingConn<C> {
1224    inner: C,
1225    read: usize,
1226    written: usize,
1227}
1228
1229impl<C> CountingConn<C> {
1230    fn new(inner: C) -> Self {
1231        CountingConn {
1232            inner,
1233            read: 0,
1234            written: 0,
1235        }
1236    }
1237}
1238
1239impl<C> AsyncRead for CountingConn<C>
1240where
1241    C: AsyncRead + Unpin,
1242{
1243    fn poll_read(
1244        self: Pin<&mut Self>,
1245        cx: &mut std::task::Context<'_>,
1246        buf: &mut io::ReadBuf<'_>,
1247    ) -> std::task::Poll<std::io::Result<()>> {
1248        let counter = self.get_mut();
1249        let pin = Pin::new(&mut counter.inner);
1250        let bytes = buf.filled().len();
1251        let poll = pin.poll_read(cx, buf);
1252        let bytes = buf.filled().len() - bytes;
1253        if let std::task::Poll::Ready(Ok(())) = poll {
1254            counter.read += bytes
1255        }
1256        poll
1257    }
1258}
1259
1260impl<C> AsyncWrite for CountingConn<C>
1261where
1262    C: AsyncWrite + Unpin,
1263{
1264    fn poll_write(
1265        self: Pin<&mut Self>,
1266        cx: &mut std::task::Context<'_>,
1267        buf: &[u8],
1268    ) -> std::task::Poll<Result<usize, std::io::Error>> {
1269        let counter = self.get_mut();
1270        let pin = Pin::new(&mut counter.inner);
1271        let poll = pin.poll_write(cx, buf);
1272        if let std::task::Poll::Ready(Ok(bytes)) = poll {
1273            counter.written += bytes
1274        }
1275        poll
1276    }
1277
1278    fn poll_flush(
1279        self: Pin<&mut Self>,
1280        cx: &mut std::task::Context<'_>,
1281    ) -> std::task::Poll<Result<(), std::io::Error>> {
1282        let counter = self.get_mut();
1283        let pin = Pin::new(&mut counter.inner);
1284        pin.poll_flush(cx)
1285    }
1286
1287    fn poll_shutdown(
1288        self: Pin<&mut Self>,
1289        cx: &mut std::task::Context<'_>,
1290    ) -> std::task::Poll<Result<(), std::io::Error>> {
1291        let counter = self.get_mut();
1292        let pin = Pin::new(&mut counter.inner);
1293        pin.poll_shutdown(cx)
1294    }
1295}
1296
1297/// Broadcasts cancellation to all matching environmentds. `conn_id`'s bits [31..20] are the lower
1298/// 12 bits of a UUID for an environmentd/organization. Using that and the template in
1299/// `cancellation_resolver` we generate a hostname. That hostname resolves to all IPs of envds that
1300/// match the UUID (cloud k8s infrastructure maintains that mapping). This function creates a new
1301/// task for each envd and relays the cancellation message to it, broadcasting it to any envd that
1302/// might match the connection.
1303///
1304/// This function returns after it has spawned the tasks, and does not wait for them to complete.
1305/// This is acceptable because cancellation in the Postgres protocol is best effort and has no
1306/// guarantees.
1307///
1308/// The safety of broadcasting this is due to the various randomness in the connection id and secret
1309/// key, which must match exactly in order to execute a query cancellation. The connection id has 19
1310/// bits of randomness, and the secret key the full 32, for a total of 51 bits. That is more than
1311/// 2e15 combinations, enough to nearly certainly prevent two different envds generating identical
1312/// combinations.
1313async fn cancel_request(
1314    conn_id: u32,
1315    secret_key: u32,
1316    cancellation_resolver: &CancellationResolver,
1317) {
1318    let suffix = conn_id_org_uuid(conn_id);
1319    let contents = match cancellation_resolver {
1320        CancellationResolver::Directory(dir) => {
1321            let path = dir.join(&suffix);
1322            match std::fs::read_to_string(&path) {
1323                Ok(contents) => contents,
1324                Err(err) => {
1325                    error!("could not read cancel file {path:?}: {err}");
1326                    return;
1327                }
1328            }
1329        }
1330        CancellationResolver::Static(addr) => addr.to_owned(),
1331    };
1332    let mut all_ips = Vec::new();
1333    for addr in contents.lines() {
1334        let addr = addr.trim();
1335        if addr.is_empty() {
1336            continue;
1337        }
1338        match tokio::net::lookup_host(addr).await {
1339            Ok(ips) => all_ips.extend(ips),
1340            Err(err) => {
1341                error!("{addr} failed resolution: {err}");
1342            }
1343        }
1344    }
1345    let mut buf = BytesMut::with_capacity(16);
1346    let msg = FrontendStartupMessage::CancelRequest {
1347        conn_id,
1348        secret_key,
1349    };
1350    msg.encode(&mut buf).expect("must encode");
1351    let buf = buf.freeze();
1352    for ip in all_ips {
1353        debug!("cancelling {suffix} to {ip}");
1354        let buf = buf.clone();
1355        spawn(|| "cancel request for ip", async move {
1356            let send = async {
1357                let mut stream = TcpStream::connect(&ip).await?;
1358                stream.write_all(&buf).await?;
1359                stream.shutdown().await?;
1360                Ok::<_, io::Error>(())
1361            };
1362            if let Err(err) = send.await {
1363                error!("error mirroring cancel to {ip}: {err}");
1364            }
1365        });
1366    }
1367}
1368
1369/// Writes an HTTP error response to a client and closes the connection.
1370///
1371/// The proxied connections are raw TCP streams, so the HTTP framing has to be written by hand.
1372/// Errors are ignored: the connection is going away regardless.
1373async fn send_http_error(client_stream: &mut Box<dyn ClientStream>, status: &str, body: &str) {
1374    let response = format!(
1375        "HTTP/1.1 {status}\r\n\
1376         Content-Type: text/plain\r\n\
1377         Content-Length: {}\r\n\
1378         Connection: close\r\n\
1379         \r\n\
1380         {body}",
1381        body.len(),
1382    );
1383    let _ = client_stream.write_all(response.as_bytes()).await;
1384    let _ = client_stream.shutdown().await;
1385}
1386
1387struct HttpsBalancer {
1388    resolver: Arc<TenantDnsResolver>,
1389    tls: Option<ReloadingSslContext>,
1390    resolve_template: Arc<str>,
1391    port: u16,
1392    metrics: Arc<ServerMetrics>,
1393    limiter: Arc<ConnectionLimiter>,
1394    configs: ConfigSet,
1395    internal_tls: Option<SslConnector>,
1396}
1397
1398impl HttpsBalancer {
1399    /// Negotiates TLS and resolves a backend using the client's SNI hostname.
1400    async fn pre_resolve(
1401        conn: Connection,
1402        tls_context: Option<ReloadingSslContext>,
1403        resolver: &TenantDnsResolver,
1404        resolve_template: &str,
1405        port: u16,
1406    ) -> Result<(Box<dyn ClientStream>, ResolvedAddr, SocketAddr), anyhow::Error> {
1407        let peer_addr = conn.peer_addr().context("fetching peer addr")?;
1408        let (client_stream, servername): (Box<dyn ClientStream>, Option<String>) = match tls_context
1409        {
1410            Some(tls_context) => {
1411                let mut ssl_stream = SslStream::new(Ssl::new(&tls_context.get())?, conn)?;
1412                if let Err(e) = Pin::new(&mut ssl_stream).accept().await {
1413                    let _ = ssl_stream.get_mut().shutdown().await;
1414                    return Err(e.into());
1415                }
1416                let servername: Option<String> =
1417                    ssl_stream.ssl().servername(NameType::HOST_NAME).map(|sn| {
1418                        match sn.split_once('.') {
1419                            Some((left, _right)) => left,
1420                            None => sn,
1421                        }
1422                        .into()
1423                    });
1424                debug!("Found sni servername: {servername:?} (https)");
1425                (Box::new(ssl_stream), servername)
1426            }
1427            _ => (Box::new(conn), None),
1428        };
1429        let resolved =
1430            Self::resolve(resolver, resolve_template, port, servername.as_deref()).await?;
1431        Ok((client_stream, resolved, peer_addr))
1432    }
1433
1434    async fn resolve(
1435        resolver: &TenantDnsResolver,
1436        resolve_template: &str,
1437        port: u16,
1438        servername: Option<&str>,
1439    ) -> Result<ResolvedAddr, anyhow::Error> {
1440        let (addr, tenant) = match servername {
1441            Some(sni) => resolver.resolve_sni(resolve_template, port, sni).await?,
1442            None => {
1443                // Without SNI, resolve the template as is. Not expected for
1444                // HTTPS in practice.
1445                debug!("https hostname (no SNI): {}:{}", resolve_template, port);
1446                resolver.resolve(resolve_template, port).await?
1447            }
1448        };
1449
1450        Ok(ResolvedAddr {
1451            addr,
1452            password: None,
1453            tenant,
1454        })
1455    }
1456}
1457
1458/// Extracts the tenant ID from an environmentd CNAME target.
1459///
1460/// The CNAME points at the environmentd service, of the form
1461/// `<service>.<namespace>.svc.cluster.local`, e.g.
1462/// `environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local`.
1463/// The `<namespace>` is the environment name `environment-<tenant_id>-<index>`,
1464/// where `<tenant_id>` is the tenant's UUID and `<index>` is the environment
1465/// generation.
1466///
1467/// NOTE: `<index>` is currently always 0, since a tenant has one environment
1468/// per region, but this does not rely on that so that multiple environments
1469/// per tenant can be supported later.
1470fn extract_tenant_from_cname(cname: &str) -> Option<String> {
1471    let mut parts = cname.split('.');
1472    let _service = parts.next();
1473    let Some(namespace) = parts.next() else {
1474        return None;
1475    };
1476    // Trim off the starting `environment-`.
1477    let Some((_, namespace)) = namespace.split_once('-') else {
1478        return None;
1479    };
1480    // Trim off the ending `-<index>`.
1481    let Some((tenant, _)) = namespace.rsplit_once('-') else {
1482        return None;
1483    };
1484    // Convert to a Uuid so that this tenant matches the frontegg resolver exactly, because it
1485    // also uses Uuid::to_string.
1486    let Ok(tenant) = Uuid::parse_str(tenant) else {
1487        error!("cname tenant not a uuid: {tenant}");
1488        return None;
1489    };
1490    Some(tenant.to_string())
1491}
1492
1493/// Strips the surrounding brackets from an IPv6 host literal, e.g. `[::1]`
1494/// becomes `::1`. Leaves other hosts unchanged. This lets a bracketed IPv6
1495/// literal in an address template parse as an `IpAddr` and resolve, matching
1496/// what `tokio::net::lookup_host` accepts.
1497fn strip_ipv6_brackets(host: &str) -> &str {
1498    host.strip_prefix('[')
1499        .and_then(|h| h.strip_suffix(']'))
1500        .unwrap_or(host)
1501}
1502
1503impl mz_server_core::Server for HttpsBalancer {
1504    const NAME: &'static str = "https_balancer";
1505
1506    // TODO(jkosh44) consider forwarding the connection UUID to the adapter.
1507    fn handle_connection(
1508        &self,
1509        conn: Connection,
1510        _tokio_metrics_intervals: impl Iterator<Item = TaskMetrics> + Send + 'static,
1511    ) -> mz_server_core::ConnectionHandler {
1512        let tls_context = self.tls.clone();
1513        let internal_tls = self.internal_tls.clone();
1514        let resolver = Arc::clone(&self.resolver);
1515        let resolve_template = Arc::clone(&self.resolve_template);
1516        let port = self.port;
1517        let inner_metrics = Arc::clone(&self.metrics);
1518        let outer_metrics = Arc::clone(&self.metrics);
1519        let limiter = Arc::clone(&self.limiter);
1520        let pre_resolved_timeout = PRE_RESOLVED_TIMEOUT.get(&self.configs);
1521        let inject_proxy_headers = INJECT_PROXY_PROTOCOL_HEADER_HTTP.get(&self.configs);
1522        Box::pin(async move {
1523            let active_guard = inner_metrics.active_connections();
1524            let result: Result<_, anyhow::Error> = Box::pin(async move {
1525                let Some(_conn_guard) = limiter.acquire() else {
1526                    return Ok(());
1527                };
1528                let pre_resolved = inner_metrics.pre_resolved_connections();
1529
1530                let (mut client_stream, resolved, peer_addr) = under_pre_resolved_timeout(
1531                    pre_resolved_timeout,
1532                    &inner_metrics,
1533                    Self::pre_resolve(conn, tls_context, &resolver, &resolve_template, port),
1534                )
1535                .await?;
1536                drop(pre_resolved);
1537                let inner_active_guard = resolved
1538                    .tenant
1539                    .as_ref()
1540                    .map(|tenant| inner_metrics.tenant_connections(tenant));
1541                let mut mz_stream = match TcpStream::connect(resolved.addr).await {
1542                    Ok(stream) => stream,
1543                    Err(e) => {
1544                        error!("failed to connect to upstream server: {e}");
1545                        send_http_error(
1546                            &mut client_stream,
1547                            "502 Bad Gateway",
1548                            "upstream server not available",
1549                        )
1550                        .await;
1551                        return Ok(());
1552                    }
1553                };
1554
1555                if inject_proxy_headers {
1556                    // Write the tcp proxy header
1557                    let addrs = ProxiedAddress::stream(peer_addr, resolved.addr);
1558                    let header = ProxyHeader::with_address(addrs);
1559                    let mut buf = [0u8; 1024];
1560                    let len = header.encode_to_slice_v2(&mut buf)?;
1561                    mz_stream.write_all(&buf[..len]).await?;
1562                }
1563
1564                let mut mz_stream = if let Some(internal_tls) = &internal_tls {
1565                    // do a TLS handshake
1566                    let mut ssl = internal_tls
1567                        .configure()?
1568                        .into_ssl(&resolved.addr.to_string())?;
1569                    ssl.set_connect_state();
1570                    Conn::Ssl(SslStream::new(ssl, mz_stream)?)
1571                } else {
1572                    Conn::Unencrypted(mz_stream)
1573                };
1574
1575                let mut client_counter = CountingConn::new(client_stream);
1576
1577                // Now blindly shuffle bytes back and forth until closed.
1578                // TODO: Limit total memory use.
1579                // See corresponding comment in pgwire implementation about ignoring the error.
1580                let _ = tokio::io::copy_bidirectional(&mut client_counter, &mut mz_stream).await;
1581                if let Some(tenant) = &resolved.tenant {
1582                    inner_metrics
1583                        .tenant_connections_tx(tenant)
1584                        .inc_by(u64::cast_from(client_counter.written));
1585                    inner_metrics
1586                        .tenant_connections_rx(tenant)
1587                        .inc_by(u64::cast_from(client_counter.read));
1588                }
1589                drop(inner_active_guard);
1590                Ok(())
1591            })
1592            .await;
1593            drop(active_guard);
1594            outer_metrics.connection_status(result.is_ok()).inc();
1595            if let Err(e) = result {
1596                debug!("connection error: {e}");
1597            }
1598            Ok(())
1599        })
1600    }
1601}
1602
1603/// Template for constructing the destination hostname from a TLS SNI
1604/// servername. `{}` is replaced with the first label of the servername.
1605#[derive(Debug)]
1606pub struct SniTemplate {
1607    pub template: String,
1608    pub port: u16,
1609}
1610
1611trait ClientStream: AsyncRead + AsyncWrite + Unpin + Send {}
1612impl<T: AsyncRead + AsyncWrite + Unpin + Send> ClientStream for T {}
1613
1614#[derive(Debug)]
1615pub enum BalancerResolver {
1616    Static(String),
1617    MultiTenant {
1618        dns: Arc<TenantDnsResolver>,
1619        frontegg: FronteggResolver,
1620        sni: Option<SniTemplate>,
1621    },
1622}
1623
1624/// An error resolving a connection's destination.
1625///
1626/// The `Display` of this error is sent to unauthenticated clients, so it must
1627/// not contain internal details such as hostnames. Those belong in the source
1628/// error attached to each variant, which is only logged.
1629#[derive(Debug, thiserror::Error)]
1630enum ResolveError {
1631    #[error("invalid password")]
1632    InvalidPassword,
1633    /// A client protocol violation, e.g. sending the wrong message during
1634    /// startup. An unauthenticated client can trigger these at will, so they
1635    /// are logged at `warn!`, not `error!`, to avoid making client noise look
1636    /// like server faults and spamming the error log.
1637    #[error("internal error")]
1638    Client(#[source] anyhow::Error),
1639    /// The tenant's upstream backend could not be reached, e.g. its hostname
1640    /// did not resolve. Reported to the client with a distinct, non-leaking
1641    /// message rather than a generic internal error, so a down environment is
1642    /// not mistaken for a balancerd bug. Logged at `warn!`: a bogus SNI reaches
1643    /// this from an unauthenticated client, and a genuinely down environment is
1644    /// an operational condition, not a balancerd fault.
1645    #[error("upstream server not available")]
1646    Upstream(#[source] anyhow::Error),
1647    /// A server-side fault.
1648    #[error("internal error")]
1649    Internal(#[from] anyhow::Error),
1650}
1651
1652impl From<io::Error> for ResolveError {
1653    fn from(e: io::Error) -> Self {
1654        ResolveError::Internal(e.into())
1655    }
1656}
1657
1658impl BalancerResolver {
1659    /// Returns a clone of the shared DNS resolver if in multi-tenant mode.
1660    /// This allows sharing the resolver with other components like HttpsBalancer.
1661    pub fn shared_dns(&self) -> Option<Arc<TenantDnsResolver>> {
1662        match self {
1663            BalancerResolver::Static(_) => None,
1664            BalancerResolver::MultiTenant { dns, .. } => Some(Arc::clone(dns)),
1665        }
1666    }
1667
1668    async fn resolve<A>(
1669        &self,
1670        conn: &mut FramedConn<A>,
1671        user: &str,
1672        metrics: &ServerMetrics,
1673    ) -> Result<ResolvedAddr, ResolveError>
1674    where
1675        A: AsyncRead + AsyncWrite + Unpin,
1676    {
1677        match self {
1678            BalancerResolver::MultiTenant {
1679                dns: dns_resolver,
1680                frontegg:
1681                    FronteggResolver {
1682                        auth,
1683                        addr_template,
1684                    },
1685                sni: sni_resolver,
1686            } => {
1687                let servername = match conn.inner() {
1688                    Conn::Ssl(ssl_stream) => {
1689                        ssl_stream.ssl().servername(NameType::HOST_NAME).map(|sn| {
1690                            match sn.split_once('.') {
1691                                Some((left, _right)) => left,
1692                                None => sn,
1693                            }
1694                        })
1695                    }
1696                    Conn::Unencrypted(_) => None,
1697                };
1698                let has_sni = servername.is_some();
1699                let resolved_addr = match (servername, sni_resolver.as_ref()) {
1700                    (Some(servername), Some(SniTemplate { template, port })) => {
1701                        // A resolution failure here means the tenant's backend
1702                        // is unreachable (or the client sent a bogus SNI). Not
1703                        // a server fault.
1704                        let (addr, tenant) = dns_resolver
1705                            .resolve_sni(template, *port, servername)
1706                            .await
1707                            .map_err(ResolveError::Upstream)?;
1708                        debug!("pgwire SNI resolved tenant: {:?}", tenant);
1709                        ResolvedAddr {
1710                            addr,
1711                            password: None,
1712                            tenant,
1713                        }
1714                    }
1715                    _ => {
1716                        conn.send(BackendMessage::AuthenticationCleartextPassword)
1717                            .await?;
1718                        conn.flush().await?;
1719                        let password = match conn.recv().await? {
1720                            Some(FrontendMessage::Password { password }) => password,
1721                            _ => {
1722                                return Err(ResolveError::Client(anyhow::anyhow!(
1723                                    "expected Password message"
1724                                )));
1725                            }
1726                        };
1727
1728                        // balancerd only needs the validated tenant_id to route
1729                        // the connection; group extraction happens in
1730                        // environmentd, so skip it here.
1731                        let auth_response = auth.authenticate(user, &password, None).await;
1732                        let auth_session = match auth_response {
1733                            Ok((auth_session, _)) => auth_session,
1734                            Err(e) => {
1735                                warn!("pgwire connection failed authentication: {}", e);
1736                                return Err(ResolveError::InvalidPassword);
1737                            }
1738                        };
1739
1740                        let hostname_with_port =
1741                            addr_template.replace("{}", &auth_session.tenant_id().to_string());
1742                        let (hostname, port_str) = hostname_with_port
1743                            .rsplit_once(':')
1744                            .ok_or_else(|| anyhow::anyhow!("port required in addr_template"))?;
1745                        let port: u16 = port_str.parse().with_context(|| {
1746                            format!("invalid port in addr_template: {}", port_str)
1747                        })?;
1748                        // The tenant is already known from authentication, so
1749                        // skip the CNAME lookup that resolve() would do. A
1750                        // failure here means the tenant's backend is unreachable.
1751                        let addr = dns_resolver
1752                            .resolve_addr(hostname, port)
1753                            .await
1754                            .map_err(ResolveError::Upstream)?;
1755                        let tenant = auth_session.tenant_id().to_string();
1756                        debug!("Frontegg resolved tenant: {}", tenant);
1757                        ResolvedAddr {
1758                            addr,
1759                            password: Some(password),
1760                            tenant: Some(tenant),
1761                        }
1762                    }
1763                };
1764                metrics
1765                    .tenant_pgwire_sni_count(
1766                        resolved_addr.tenant.as_deref().unwrap_or("unknown"),
1767                        has_sni,
1768                    )
1769                    .inc();
1770
1771                Ok(resolved_addr)
1772            }
1773            BalancerResolver::Static(addr) => {
1774                // We don't want any caching here so we just use the standard
1775                // tokio resolver.
1776                let Some(addr) = tokio::net::lookup_host(addr).await?.next() else {
1777                    return Err(anyhow::anyhow!("{addr} did not resolve to any addresses").into());
1778                };
1779
1780                Ok(ResolvedAddr {
1781                    addr,
1782                    password: None,
1783                    tenant: None,
1784                })
1785            }
1786        }
1787    }
1788}
1789
1790/// Creates a resolver from the system DNS configuration.
1791///
1792/// Caching is delegated to the infrastructure (node-local DNS), so this
1793/// resolver does no caching of its own. Fails if the system DNS configuration
1794/// cannot be read. We must not fall back to hickory's default config (Google
1795/// public DNS) here, that would leak internal hostnames to an external party
1796/// and could not resolve them anyway.
1797fn create_resolver() -> Result<TokioResolver, anyhow::Error> {
1798    let (config, mut opts) = read_system_conf().context("reading system DNS configuration")?;
1799    opts.cache_size = 0;
1800    // Query A records first and AAAA only on failure, rather than both.
1801    opts.ip_strategy = LookupIpStrategy::Ipv4thenIpv6;
1802
1803    Ok(
1804        Resolver::builder_with_config(config, TokioRuntimeProvider::default())
1805            .with_options(opts)
1806            .build()?,
1807    )
1808}
1809
1810/// Resolves tenant hostnames for pgwire and HTTPS routing.
1811///
1812/// Caching is delegated to the infrastructure (node-local DNS), so every
1813/// lookup issues a query. CNAMEs are resolved separately from A records only
1814/// because the CNAME carries the tenant, not for caching reasons.
1815#[derive(Debug)]
1816pub struct TenantDnsResolver {
1817    resolver: TokioResolver,
1818}
1819
1820impl TenantDnsResolver {
1821    /// Creates a new resolver. Fails if the system DNS configuration cannot be
1822    /// read.
1823    pub fn new() -> Result<Self, anyhow::Error> {
1824        Ok(Self {
1825            resolver: create_resolver()?,
1826        })
1827    }
1828
1829    /// Resolves the CNAME a hostname points at, if any.
1830    async fn resolve_cname(&self, hostname: &str) -> Option<String> {
1831        match self.resolver.lookup(hostname, RecordType::CNAME).await {
1832            Ok(cname_response) => {
1833                if let Some(cname_record) = cname_response.answers().first() {
1834                    if let RData::CNAME(cname_data) = &cname_record.data {
1835                        let cname = cname_data.to_string();
1836                        debug!("CNAME for {}: {}", hostname, cname);
1837                        return Some(cname);
1838                    }
1839                }
1840                None
1841            }
1842            Err(e) => {
1843                debug!("CNAME lookup failed for {}: {}", hostname, e);
1844                None
1845            }
1846        }
1847    }
1848
1849    /// Resolves the A records for a hostname.
1850    async fn resolve_a(&self, hostname: &str) -> Result<LookupIp, anyhow::Error> {
1851        self.resolver
1852            .lookup_ip(hostname)
1853            .await
1854            .with_context(|| format!("resolving A records for {}", hostname))
1855    }
1856
1857    /// Resolves the environment address for a TLS SNI servername.
1858    ///
1859    /// `servername` is the first label of the SNI host, e.g.
1860    /// `3dl07g8zmj91pntk4eo9cfvwe`. Substituting it into `template` (e.g.
1861    /// `blncr-{}`) yields a Kubernetes hostname like
1862    /// `blncr-3dl07g8zmj91pntk4eo9cfvwe`, which resolves via a CNAME to the
1863    /// environmentd service, e.g.
1864    /// `environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local`.
1865    /// The tenant is extracted from that CNAME. See
1866    /// `extract_tenant_from_cname`.
1867    pub async fn resolve_sni(
1868        &self,
1869        template: &str,
1870        port: u16,
1871        servername: &str,
1872    ) -> Result<(SocketAddr, Option<String>), anyhow::Error> {
1873        let hostname = template.replace("{}", servername);
1874        debug!("SNI hostname: {}", hostname);
1875        self.resolve(&hostname, port).await
1876    }
1877
1878    /// Resolves the address for a hostname, skipping CNAME resolution and
1879    /// tenant extraction. Use when the tenant is already known.
1880    async fn resolve_addr(&self, host: &str, port: u16) -> Result<SocketAddr, anyhow::Error> {
1881        let host = strip_ipv6_brackets(host);
1882        // IP literals need no resolution. Resolving them through hickory
1883        // would walk the search domain list first when ndots is large, as it
1884        // is in Kubernetes.
1885        if let Ok(ip) = host.parse::<IpAddr>() {
1886            return Ok(SocketAddr::new(ip, port));
1887        }
1888        Self::first_addr(self.resolve_a(host).await?, port)
1889    }
1890
1891    /// Resolves the address and tenant from a hostname and port.
1892    ///
1893    /// The tenant is extracted from the CNAME if the hostname points at one.
1894    async fn resolve(
1895        &self,
1896        host: &str,
1897        port: u16,
1898    ) -> Result<(SocketAddr, Option<String>), anyhow::Error> {
1899        let host = strip_ipv6_brackets(host);
1900        // IP literals need no resolution and carry no tenant CNAME.
1901        if let Ok(ip) = host.parse::<IpAddr>() {
1902            return Ok((SocketAddr::new(ip, port), None));
1903        }
1904
1905        // The CNAME carries the tenant, so resolve it separately to extract it.
1906        let (ips, tenant) = if let Some(cname) = self.resolve_cname(host).await {
1907            let tenant = extract_tenant_from_cname(&cname);
1908            (self.resolve_a(&cname).await?, tenant)
1909        } else {
1910            (self.resolve_a(host).await?, None)
1911        };
1912
1913        Ok((Self::first_addr(ips, port)?, tenant))
1914    }
1915
1916    /// Returns the first resolved IP as a socket address.
1917    fn first_addr(ips: LookupIp, port: u16) -> Result<SocketAddr, anyhow::Error> {
1918        ips.iter()
1919            .next()
1920            .map(|ip| SocketAddr::new(ip, port))
1921            .ok_or_else(|| anyhow::anyhow!("no A records found in DNS response"))
1922    }
1923}
1924
1925#[derive(Debug)]
1926pub struct FronteggResolver {
1927    pub auth: FronteggAuthentication,
1928    pub addr_template: String,
1929}
1930
1931#[derive(Debug)]
1932struct ResolvedAddr {
1933    addr: SocketAddr,
1934    password: Option<String>,
1935    tenant: Option<String>,
1936}
1937
1938#[cfg(test)]
1939mod tests {
1940    use mz_dyncfg::ConfigUpdates;
1941
1942    use super::*;
1943
1944    #[mz_ore::test]
1945    fn test_tenant() {
1946        let tests = vec![
1947            ("", None),
1948            (
1949                "environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local",
1950                Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1951            ),
1952            (
1953                // Trailing dot from an absolute DNS name, as returned by the
1954                // resolver.
1955                "environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local.",
1956                Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1957            ),
1958            (
1959                // Variously named parts.
1960                "service.something-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.ssvvcc.cloister.faraway",
1961                Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1962            ),
1963            (
1964                // No dashes in uuid.
1965                "environmentd.environment-58cd23ffa4d74bd0ad85a6ff29cc86c3-0.svc.cluster.local",
1966                Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1967            ),
1968            (
1969                // -1234 suffix.
1970                "environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-1234.svc.cluster.local",
1971                Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1972            ),
1973            (
1974                // Uppercase.
1975                "environmentd.environment-58CD23FF-A4D7-4BD0-AD85-A6FF29CC86C3-0.svc.cluster.local",
1976                Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1977            ),
1978            (
1979                // No -number suffix.
1980                "environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3.svc.cluster.local",
1981                None,
1982            ),
1983            (
1984                // No service name.
1985                "environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local",
1986                None,
1987            ),
1988            (
1989                // Invalid UUID.
1990                "environmentd.environment-8cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local",
1991                None,
1992            ),
1993        ];
1994        for (name, expect) in tests {
1995            let cname = extract_tenant_from_cname(name);
1996            assert_eq!(
1997                cname.as_deref(),
1998                expect,
1999                "{name} got {cname:?} expected {expect:?}"
2000            );
2001        }
2002    }
2003
2004    #[mz_ore::test]
2005    fn test_strip_ipv6_brackets() {
2006        assert_eq!(strip_ipv6_brackets("[::1]"), "::1");
2007        assert_eq!(strip_ipv6_brackets("[2001:db8::1]"), "2001:db8::1");
2008        // Hosts without a matched bracket pair are left untouched.
2009        assert_eq!(strip_ipv6_brackets("127.0.0.1"), "127.0.0.1");
2010        assert_eq!(strip_ipv6_brackets("host.example.com"), "host.example.com");
2011        assert_eq!(strip_ipv6_brackets("[unclosed"), "[unclosed");
2012        assert_eq!(strip_ipv6_brackets("unopened]"), "unopened]");
2013    }
2014
2015    #[mz_ore::test]
2016    fn test_connection_limiter() {
2017        let configs = dyncfgs::all_dyncfgs(ConfigSet::default());
2018        let set_max = |max: u32| {
2019            let mut updates = ConfigUpdates::default();
2020            updates.add(&MAX_CONNECTIONS, max);
2021            updates.apply(&configs);
2022        };
2023
2024        set_max(2);
2025        let limiter = ConnectionLimiter::new(&MetricsRegistry::new(), configs.clone());
2026        let first = limiter.acquire().expect("under the limit");
2027        let second = limiter.acquire().expect("at the limit");
2028        assert!(limiter.acquire().is_none());
2029        assert_eq!(limiter.rejected.get(), 1);
2030
2031        // A connection closing frees capacity for a new one.
2032        drop(second);
2033        let _third = limiter.acquire().expect("capacity freed");
2034        drop(first);
2035
2036        // Lowering the limit below the current count refuses new connections but leaves the
2037        // established ones alone.
2038        set_max(1);
2039        assert!(limiter.acquire().is_none());
2040
2041        // Zero disables the limit.
2042        set_max(0);
2043        assert!(limiter.acquire().is_some());
2044    }
2045
2046    /// A zero timeout disables the deadline, which is what leaves the phase unbounded.
2047    #[mz_ore::test(tokio::test(start_paused = true))]
2048    async fn test_pre_resolved_timeout_zero_disables() {
2049        let metrics = ServerMetrics::new(
2050            ServerMetricsConfig::register_into(&MetricsRegistry::new()),
2051            "pgwire",
2052        );
2053
2054        // Zero lets the phase run as long as it likes. The paused clock makes the hour free.
2055        let slow = async {
2056            tokio::time::sleep(Duration::from_secs(60 * 60)).await;
2057            Ok::<(), io::Error>(())
2058        };
2059        assert!(
2060            under_pre_resolved_timeout(Duration::ZERO, &metrics, slow)
2061                .await
2062                .is_ok()
2063        );
2064        assert_eq!(metrics.pre_resolved_timeout().get(), 0);
2065
2066        // Anything else bounds it, and the overrun is counted.
2067        let stalled = std::future::pending::<Result<(), io::Error>>();
2068        let err = under_pre_resolved_timeout(Duration::from_secs(30), &metrics, stalled)
2069            .await
2070            .expect_err("a phase that never finishes must be cut short");
2071        assert_eq!(err.kind(), io::ErrorKind::TimedOut);
2072        assert_eq!(metrics.pre_resolved_timeout().get(), 1);
2073    }
2074}