1mod 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
87pub const BUILD_INFO: BuildInfo = build_info!();
89
90pub struct BalancerConfig {
91 build_version: Version,
93 internal_http_listen_addr: SocketAddr,
95 pgwire_listen_addr: SocketAddr,
97 https_listen_addr: SocketAddr,
99 cancellation_resolver: CancellationResolver,
101 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#[derive(Debug)]
166pub struct BalancerMetrics {
167 _uptime: ComputedGauge,
168}
169
170impl BalancerMetrics {
171 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 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) .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 .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 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 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 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 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 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
497fn internal_tls_connector() -> Result<SslConnector, ErrorStack> {
499 let mut builder = SslConnector::builder(SslMethod::tls())?;
500 builder.set_verify(SslVerifyMode::NONE);
502 Ok(builder.build())
503}
504
505struct 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 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#[derive(Debug)]
674struct ConnectionLimiter {
675 configs: ConfigSet,
676 active: AtomicU32,
677 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 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 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
743struct 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
752async 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
766async 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 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 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 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 ¶ms,
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 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 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 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 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 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 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 if password.is_none() {
1122 return Ok(mz_stream);
1123 }
1124
1125 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 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 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 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 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
1222struct 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
1297async 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
1369async 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 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 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
1458fn 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 let Some((_, namespace)) = namespace.split_once('-') else {
1478 return None;
1479 };
1480 let Some((tenant, _)) = namespace.rsplit_once('-') else {
1482 return None;
1483 };
1484 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
1493fn 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 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 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 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 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#[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#[derive(Debug, thiserror::Error)]
1630enum ResolveError {
1631 #[error("invalid password")]
1632 InvalidPassword,
1633 #[error("internal error")]
1638 Client(#[source] anyhow::Error),
1639 #[error("upstream server not available")]
1646 Upstream(#[source] anyhow::Error),
1647 #[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 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 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 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 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 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
1790fn 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 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#[derive(Debug)]
1816pub struct TenantDnsResolver {
1817 resolver: TokioResolver,
1818}
1819
1820impl TenantDnsResolver {
1821 pub fn new() -> Result<Self, anyhow::Error> {
1824 Ok(Self {
1825 resolver: create_resolver()?,
1826 })
1827 }
1828
1829 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 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 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 async fn resolve_addr(&self, host: &str, port: u16) -> Result<SocketAddr, anyhow::Error> {
1881 let host = strip_ipv6_brackets(host);
1882 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 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 if let Ok(ip) = host.parse::<IpAddr>() {
1902 return Ok((SocketAddr::new(ip, port), None));
1903 }
1904
1905 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 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 "environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local.",
1956 Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1957 ),
1958 (
1959 "service.something-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.ssvvcc.cloister.faraway",
1961 Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1962 ),
1963 (
1964 "environmentd.environment-58cd23ffa4d74bd0ad85a6ff29cc86c3-0.svc.cluster.local",
1966 Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1967 ),
1968 (
1969 "environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-1234.svc.cluster.local",
1971 Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1972 ),
1973 (
1974 "environmentd.environment-58CD23FF-A4D7-4BD0-AD85-A6FF29CC86C3-0.svc.cluster.local",
1976 Some("58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3"),
1977 ),
1978 (
1979 "environmentd.environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3.svc.cluster.local",
1981 None,
1982 ),
1983 (
1984 "environment-58cd23ff-a4d7-4bd0-ad85-a6ff29cc86c3-0.svc.cluster.local",
1986 None,
1987 ),
1988 (
1989 "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 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 drop(second);
2033 let _third = limiter.acquire().expect("capacity freed");
2034 drop(first);
2035
2036 set_max(1);
2039 assert!(limiter.acquire().is_none());
2040
2041 set_max(0);
2043 assert!(limiter.acquire().is_some());
2044 }
2045
2046 #[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 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 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}