Skip to main content

mz_environmentd/
test_util.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
10use std::collections::BTreeMap;
11use std::error::Error;
12use std::future::IntoFuture;
13use std::net::{IpAddr, Ipv4Addr, SocketAddr, TcpStream};
14use std::path::{Path, PathBuf};
15use std::pin::Pin;
16use std::str::FromStr;
17use std::sync::Arc;
18use std::sync::LazyLock;
19use std::time::Duration;
20use std::{env, fs, iter};
21
22use anyhow::anyhow;
23use futures::Future;
24use futures::future::{BoxFuture, LocalBoxFuture};
25use headers::{Header, HeaderMapExt};
26use http::Uri;
27use hyper::http::header::HeaderMap;
28use maplit::btreemap;
29use mz_adapter::TimestampExplanation;
30use mz_adapter_types::bootstrap_builtin_cluster_config::{
31    ANALYTICS_CLUSTER_DEFAULT_REPLICATION_FACTOR, BootstrapBuiltinClusterConfig,
32    CATALOG_SERVER_CLUSTER_DEFAULT_REPLICATION_FACTOR, PROBE_CLUSTER_DEFAULT_REPLICATION_FACTOR,
33    SUPPORT_CLUSTER_DEFAULT_REPLICATION_FACTOR, SYSTEM_CLUSTER_DEFAULT_REPLICATION_FACTOR,
34};
35use mz_adapter_types::dyncfgs::ENABLE_CLUSTER_RECONFIGURATION_LAG_GATE;
36
37use mz_auth::password::Password;
38use mz_catalog::config::ClusterReplicaSizeMap;
39use mz_controller::ControllerConfig;
40use mz_dyncfg::ConfigUpdates;
41use mz_license_keys::ValidatedLicenseKey;
42use mz_orchestrator_process::{ProcessOrchestrator, ProcessOrchestratorConfig};
43use mz_orchestrator_tracing::{TracingCliArgs, TracingOrchestrator};
44use mz_ore::cast::CastLossy;
45use mz_ore::metrics::MetricsRegistry;
46use mz_ore::now::{EpochMillis, NowFn, SYSTEM_TIME};
47use mz_ore::retry::Retry;
48use mz_ore::task;
49use mz_ore::tracing::{
50    OpenTelemetryConfig, StderrLogConfig, StderrLogFormat, TracingConfig, TracingHandle,
51};
52use mz_persist_client::PersistLocation;
53use mz_persist_client::cache::PersistClientCache;
54use mz_persist_client::cfg::{CONSENSUS_CONNECTION_POOL_MAX_SIZE, PersistConfig};
55use mz_persist_client::rpc::PersistGrpcPubSubServer;
56use mz_postgres_util::{
57    Sql, batch_execute as pg_batch_execute, execute as pg_execute, query_one as pg_query_one, sql,
58};
59use mz_secrets::SecretsController;
60use mz_server_core::listeners::v26_32_0::ListenersConfig;
61use mz_server_core::listeners::{
62    AllowedRoles, AuthenticatorKind, HttpListenerConfig, HttpRoutesEnabled, RouteGroup,
63};
64use mz_server_core::{ReloadTrigger, TlsCertConfig};
65use mz_sql::catalog::EnvironmentId;
66use mz_storage_types::connections::ConnectionContext;
67use mz_tracing::CloneableEnvFilter;
68use openssl::asn1::Asn1Time;
69use openssl::error::ErrorStack;
70use openssl::hash::MessageDigest;
71use openssl::nid::Nid;
72use openssl::pkey::{PKey, Private};
73use openssl::rsa::Rsa;
74use openssl::ssl::{SslConnector, SslConnectorBuilder, SslMethod, SslOptions};
75use openssl::x509::extension::{BasicConstraints, SubjectAlternativeName};
76use openssl::x509::{X509, X509Name, X509NameBuilder};
77use postgres::error::DbError;
78use postgres::tls::{MakeTlsConnect, TlsConnect};
79use postgres::types::{FromSql, Type};
80use postgres::{NoTls, Socket};
81use postgres_openssl::MakeTlsConnector;
82use tempfile::TempDir;
83use tokio::net::TcpListener;
84use tokio::runtime::Runtime;
85use tokio_postgres::config::{Host, SslMode};
86use tokio_postgres::{AsyncMessage, Client};
87use tokio_stream::wrappers::TcpListenerStream;
88use tower_http::cors::AllowOrigin;
89use tracing::Level;
90use tracing_capture::SharedStorage;
91use tracing_subscriber::EnvFilter;
92use tungstenite::stream::MaybeTlsStream;
93use tungstenite::{Message, WebSocket};
94
95use crate::{
96    CatalogConfig, FronteggAuthenticator, SqlListenerConfig, WebSocketAuth, WebSocketResponse,
97};
98
99pub static KAFKA_ADDRS: LazyLock<String> =
100    LazyLock::new(|| env::var("KAFKA_ADDRS").unwrap_or_else(|_| "localhost:9092".into()));
101
102/// Entry point for creating and configuring an `environmentd` test harness.
103#[derive(Clone)]
104pub struct TestHarness {
105    data_directory: Option<PathBuf>,
106    tls: Option<TlsCertConfig>,
107    frontegg: Option<FronteggAuthenticator>,
108    external_login_password_mz_system: Option<Password>,
109    listeners_config: ListenersConfig,
110    unsafe_mode: bool,
111    /// Whether the connection context carries the AWS external ID prefix and
112    /// connection role ARN. Default true, matching a cloud deployment. Set false
113    /// via [`TestHarness::without_aws_connection_context`] to model a deployment
114    /// that never configured the AWS context, where the context functions fold
115    /// to NULL.
116    aws_connection_context: bool,
117    workers: usize,
118    now: NowFn,
119    seed: u32,
120    storage_usage_collection_interval: Duration,
121    storage_usage_retention_period: Option<Duration>,
122    default_cluster_replica_size: String,
123    default_cluster_replication_factor: u32,
124    builtin_system_cluster_config: BootstrapBuiltinClusterConfig,
125    builtin_catalog_server_cluster_config: BootstrapBuiltinClusterConfig,
126    builtin_probe_cluster_config: BootstrapBuiltinClusterConfig,
127    builtin_support_cluster_config: BootstrapBuiltinClusterConfig,
128    builtin_analytics_cluster_config: BootstrapBuiltinClusterConfig,
129
130    propagate_crashes: bool,
131    enable_tracing: bool,
132    // This is currently unrelated to enable_tracing, and is used only to disable orchestrator
133    // tracing.
134    orchestrator_tracing_cli_args: TracingCliArgs,
135    bootstrap_role: Option<String>,
136    deploy_generation: u64,
137    system_parameter_defaults: BTreeMap<String, String>,
138    internal_console_redirect_url: Option<String>,
139    metrics_registry: Option<MetricsRegistry>,
140    code_version: semver::Version,
141    force_builtin_schema_migration: Option<String>,
142    capture: Option<SharedStorage>,
143    pub environment_id: EnvironmentId,
144}
145
146impl Default for TestHarness {
147    fn default() -> TestHarness {
148        TestHarness {
149            data_directory: None,
150            tls: None,
151            frontegg: None,
152            external_login_password_mz_system: None,
153            listeners_config: ListenersConfig {
154                sql: btreemap![
155                    "external".to_owned() => SqlListenerConfig {
156                        addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
157                        authenticator_kind: AuthenticatorKind::None,
158                        allowed_roles: AllowedRoles::Normal,
159                        enable_tls: false,
160                    },
161                    "internal".to_owned() => SqlListenerConfig {
162                        addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
163                        authenticator_kind: AuthenticatorKind::None,
164                        allowed_roles: AllowedRoles::NormalAndInternal,
165                        enable_tls: false,
166                    },
167                ],
168                http: btreemap![
169                    "external".to_owned() => HttpListenerConfig {
170                        addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
171                        authenticator_kind: AuthenticatorKind::None,
172                        enable_tls: false,
173                        routes: HttpRoutesEnabled {
174                            base: RouteGroup::Enabled(AllowedRoles::Normal),
175                            webhook: RouteGroup::Enabled(AllowedRoles::Normal),
176                            internal: RouteGroup::Disabled,
177                            metrics: RouteGroup::Disabled,
178                            profiling: RouteGroup::Disabled,
179                            mcp_agent: RouteGroup::Disabled,
180                            mcp_developer: RouteGroup::Disabled,
181                            console_config: RouteGroup::Enabled(AllowedRoles::Normal),
182                        },
183                    },
184                    "internal".to_owned() => HttpListenerConfig {
185                        addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
186                        authenticator_kind: AuthenticatorKind::None,
187                        enable_tls: false,
188                        routes: HttpRoutesEnabled {
189                            base: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
190                            webhook: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
191                            internal: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
192                            metrics: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
193                            profiling: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
194                            mcp_agent: RouteGroup::Disabled,
195                            mcp_developer: RouteGroup::Disabled,
196                            console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
197                        },
198                    },
199                ],
200            },
201            unsafe_mode: false,
202            aws_connection_context: true,
203            workers: 1,
204            now: SYSTEM_TIME.clone(),
205            seed: rand::random(),
206            storage_usage_collection_interval: Duration::from_secs(3600),
207            storage_usage_retention_period: None,
208            default_cluster_replica_size: "scale=1,workers=1".to_string(),
209            default_cluster_replication_factor: 1,
210            builtin_system_cluster_config: BootstrapBuiltinClusterConfig {
211                size: "scale=1,workers=1".to_string(),
212                replication_factor: SYSTEM_CLUSTER_DEFAULT_REPLICATION_FACTOR,
213            },
214            builtin_catalog_server_cluster_config: BootstrapBuiltinClusterConfig {
215                size: "scale=1,workers=1".to_string(),
216                replication_factor: CATALOG_SERVER_CLUSTER_DEFAULT_REPLICATION_FACTOR,
217            },
218            builtin_probe_cluster_config: BootstrapBuiltinClusterConfig {
219                size: "scale=1,workers=1".to_string(),
220                replication_factor: PROBE_CLUSTER_DEFAULT_REPLICATION_FACTOR,
221            },
222            builtin_support_cluster_config: BootstrapBuiltinClusterConfig {
223                size: "scale=1,workers=1".to_string(),
224                replication_factor: SUPPORT_CLUSTER_DEFAULT_REPLICATION_FACTOR,
225            },
226            builtin_analytics_cluster_config: BootstrapBuiltinClusterConfig {
227                size: "scale=1,workers=1".to_string(),
228                replication_factor: ANALYTICS_CLUSTER_DEFAULT_REPLICATION_FACTOR,
229            },
230            propagate_crashes: false,
231            enable_tracing: false,
232            bootstrap_role: Some("materialize".into()),
233            deploy_generation: 0,
234            // This and startup_log_filter below are both (?) needed to suppress clusterd messages.
235            // If we need those in the future, we might need to change both.
236            system_parameter_defaults: BTreeMap::from([
237                ("log_filter".to_string(), "error".to_string()),
238                (
239                    ENABLE_CLUSTER_RECONFIGURATION_LAG_GATE.name().to_string(),
240                    "true".to_string(),
241                ),
242            ]),
243            internal_console_redirect_url: None,
244            metrics_registry: None,
245            orchestrator_tracing_cli_args: TracingCliArgs {
246                startup_log_filter: CloneableEnvFilter::from_str("error").expect("must parse"),
247                ..Default::default()
248            },
249            code_version: crate::BUILD_INFO.semver_version(),
250            force_builtin_schema_migration: None,
251            environment_id: EnvironmentId::for_tests(),
252            capture: None,
253        }
254    }
255}
256
257impl TestHarness {
258    /// Starts a test [`TestServer`], panicking if the server could not be started.
259    ///
260    /// For cases when startup might fail, see [`TestHarness::try_start`].
261    pub async fn start(self) -> TestServer {
262        self.try_start().await.expect("Failed to start test Server")
263    }
264
265    /// Like [`TestHarness::start`] but can specify a cert reload trigger.
266    pub async fn start_with_trigger(self, tls_reload_certs: ReloadTrigger) -> TestServer {
267        self.try_start_with_trigger(tls_reload_certs)
268            .await
269            .expect("Failed to start test Server")
270    }
271
272    /// Starts a test [`TestServer`], returning an error if the server could not be started.
273    pub async fn try_start(self) -> Result<TestServer, anyhow::Error> {
274        self.try_start_with_trigger(mz_server_core::cert_reload_never_reload())
275            .await
276    }
277
278    /// Like [`TestHarness::try_start`] but can specify a cert reload trigger.
279    pub async fn try_start_with_trigger(
280        self,
281        tls_reload_certs: ReloadTrigger,
282    ) -> Result<TestServer, anyhow::Error> {
283        let listeners = Listeners::new(&self).await?;
284        listeners.serve_with_trigger(self, tls_reload_certs).await
285    }
286
287    /// Starts a runtime and returns a [`TestServerWithRuntime`].
288    pub fn start_blocking(self) -> TestServerWithRuntime {
289        let runtime = tokio::runtime::Builder::new_multi_thread()
290            .enable_all()
291            .thread_stack_size(mz_ore::stack::STACK_SIZE)
292            .build()
293            .expect("failed to spawn runtime for test");
294        let runtime = Arc::new(runtime);
295        let server = runtime.block_on(self.start());
296        TestServerWithRuntime { runtime, server }
297    }
298
299    pub fn data_directory(mut self, data_directory: impl Into<PathBuf>) -> Self {
300        self.data_directory = Some(data_directory.into());
301        self
302    }
303
304    pub fn with_tls(mut self, cert_path: impl Into<PathBuf>, key_path: impl Into<PathBuf>) -> Self {
305        self.tls = Some(TlsCertConfig {
306            cert: cert_path.into(),
307            key: key_path.into(),
308        });
309        for (_, listener) in &mut self.listeners_config.sql {
310            listener.enable_tls = true;
311        }
312        for (_, listener) in &mut self.listeners_config.http {
313            listener.enable_tls = true;
314        }
315        self
316    }
317
318    pub fn unsafe_mode(mut self) -> Self {
319        self.unsafe_mode = true;
320        self
321    }
322
323    /// Models a deployment that never configured the AWS context, so the AWS
324    /// external ID prefix and connection role ARN are absent and the plan-time
325    /// AWS context functions fold to NULL.
326    pub fn without_aws_connection_context(mut self) -> Self {
327        self.aws_connection_context = false;
328        self
329    }
330
331    pub fn workers(mut self, workers: usize) -> Self {
332        self.workers = workers;
333        self
334    }
335
336    pub fn with_frontegg_auth(mut self, frontegg: &FronteggAuthenticator) -> Self {
337        self.frontegg = Some(frontegg.clone());
338        let enable_tls = self.tls.is_some();
339        self.listeners_config = ListenersConfig {
340            sql: btreemap! {
341                "external".to_owned() => SqlListenerConfig {
342                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
343                    authenticator_kind: AuthenticatorKind::Frontegg,
344                    allowed_roles: AllowedRoles::Normal,
345                    enable_tls,
346                },
347                "internal".to_owned() => SqlListenerConfig {
348                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
349                    authenticator_kind: AuthenticatorKind::None,
350                    allowed_roles: AllowedRoles::NormalAndInternal,
351                    enable_tls: false,
352                },
353            },
354            http: btreemap! {
355                "external".to_owned() => HttpListenerConfig {
356                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
357                    authenticator_kind: AuthenticatorKind::Frontegg,
358                    enable_tls,
359                    routes: HttpRoutesEnabled {
360                        base: RouteGroup::Enabled(AllowedRoles::Normal),
361                        webhook: RouteGroup::Enabled(AllowedRoles::Normal),
362                        internal: RouteGroup::Disabled,
363                        metrics: RouteGroup::Disabled,
364                        profiling: RouteGroup::Disabled,
365                        mcp_agent: RouteGroup::Disabled,
366                        mcp_developer: RouteGroup::Disabled,
367                        console_config: RouteGroup::Enabled(AllowedRoles::Normal),
368                    },
369                },
370                "internal".to_owned() => HttpListenerConfig {
371                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
372                    authenticator_kind: AuthenticatorKind::None,
373                    enable_tls: false,
374                    routes: HttpRoutesEnabled {
375                        base: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
376                        webhook: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
377                        internal: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
378                        metrics: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
379                        profiling: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
380                        mcp_agent: RouteGroup::Disabled,
381                        mcp_developer: RouteGroup::Disabled,
382                        console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
383                    },
384                },
385            },
386        };
387        self
388    }
389
390    pub fn with_oidc_auth(
391        mut self,
392        issuer: Option<String>,
393        authentication_claim: Option<String>,
394        expected_audiences: Option<Vec<String>>,
395        external_login_password_mz_system: Option<Password>,
396    ) -> Self {
397        let enable_tls = self.tls.is_some();
398        self.listeners_config = ListenersConfig {
399            sql: btreemap! {
400                "external".to_owned() => SqlListenerConfig {
401                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
402                    authenticator_kind: AuthenticatorKind::Oidc,
403                    allowed_roles: AllowedRoles::NormalAndInternal,
404                    enable_tls,
405                },
406                "internal".to_owned() => SqlListenerConfig {
407                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
408                    authenticator_kind: AuthenticatorKind::None,
409                    allowed_roles: AllowedRoles::NormalAndInternal,
410                    enable_tls: false,
411                },
412            },
413            http: btreemap! {
414                "external".to_owned() => HttpListenerConfig {
415                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
416                    authenticator_kind: AuthenticatorKind::Oidc,
417                    enable_tls,
418                    routes: HttpRoutesEnabled {
419                        base: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
420                        webhook: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
421                        internal: RouteGroup::Disabled,
422                        metrics: RouteGroup::Disabled,
423                        profiling: RouteGroup::Disabled,
424                        mcp_agent: RouteGroup::Disabled,
425                        mcp_developer: RouteGroup::Disabled,
426                        console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
427                    },
428                },
429                "internal".to_owned() => HttpListenerConfig {
430                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
431                    authenticator_kind: AuthenticatorKind::None,
432                    enable_tls: false,
433                    routes: HttpRoutesEnabled {
434                        base: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
435                        webhook: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
436                        internal: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
437                        metrics: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
438                        profiling: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
439                        mcp_agent: RouteGroup::Disabled,
440                        mcp_developer: RouteGroup::Disabled,
441                        console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
442                    },
443                },
444            },
445        };
446
447        if let Some(issuer) = issuer {
448            self.system_parameter_defaults
449                .insert("oidc_issuer".to_string(), issuer);
450        }
451
452        if let Some(authentication_claim) = authentication_claim {
453            self.system_parameter_defaults.insert(
454                "oidc_authentication_claim".to_string(),
455                authentication_claim,
456            );
457        }
458
459        if let Some(expected_audiences) = expected_audiences {
460            self.system_parameter_defaults.insert(
461                "oidc_audience".to_string(),
462                serde_json::to_string(&expected_audiences).unwrap(),
463            );
464        }
465
466        if let Some(external_login_password_mz_system) = external_login_password_mz_system {
467            self.external_login_password_mz_system = Some(external_login_password_mz_system);
468            self.system_parameter_defaults
469                .insert("enable_password_auth".to_string(), "true".to_string());
470        }
471
472        self
473    }
474
475    pub fn with_password_auth(mut self, mz_system_password: Password) -> Self {
476        self.external_login_password_mz_system = Some(mz_system_password);
477        let enable_tls = self.tls.is_some();
478        self.listeners_config = ListenersConfig {
479            sql: btreemap! {
480                "external".to_owned() => SqlListenerConfig {
481                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
482                    authenticator_kind: AuthenticatorKind::Password,
483                    allowed_roles: AllowedRoles::NormalAndInternal,
484                    enable_tls,
485                },
486            },
487            http: btreemap! {
488                "external".to_owned() => HttpListenerConfig {
489                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
490                    authenticator_kind: AuthenticatorKind::Password,
491                    enable_tls,
492                    routes: HttpRoutesEnabled {
493                        base: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
494                        webhook: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
495                        internal: RouteGroup::Enabled(AllowedRoles::Internal),
496                        metrics: RouteGroup::Disabled,
497                        profiling: RouteGroup::Enabled(AllowedRoles::Internal),
498                        mcp_agent: RouteGroup::Disabled,
499                        mcp_developer: RouteGroup::Disabled,
500                        console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
501                    },
502                },
503                "metrics".to_owned() => HttpListenerConfig {
504                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
505                    authenticator_kind: AuthenticatorKind::None,
506                    enable_tls: false,
507                    routes: HttpRoutesEnabled {
508                        base: RouteGroup::Disabled,
509                        webhook: RouteGroup::Disabled,
510                        internal: RouteGroup::Disabled,
511                        metrics: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
512                        profiling: RouteGroup::Disabled,
513                        mcp_agent: RouteGroup::Disabled,
514                        mcp_developer: RouteGroup::Disabled,
515                        console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
516                    },
517                },
518            },
519        };
520        self
521    }
522
523    pub fn with_sasl_scram_auth(mut self, mz_system_password: Password) -> Self {
524        self.external_login_password_mz_system = Some(mz_system_password);
525        let enable_tls = self.tls.is_some();
526        self.listeners_config = ListenersConfig {
527            sql: btreemap! {
528                "external".to_owned() => SqlListenerConfig {
529                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
530                    authenticator_kind: AuthenticatorKind::Sasl,
531                    allowed_roles: AllowedRoles::NormalAndInternal,
532                    enable_tls,
533                },
534            },
535            http: btreemap! {
536                "external".to_owned() => HttpListenerConfig {
537                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
538                    authenticator_kind: AuthenticatorKind::Password,
539                    enable_tls,
540                    routes: HttpRoutesEnabled {
541                        base: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
542                        webhook: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
543                        internal: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
544                        metrics: RouteGroup::Disabled,
545                        profiling: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
546                        mcp_agent: RouteGroup::Disabled,
547                        mcp_developer: RouteGroup::Disabled,
548                        console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
549                    },
550                },
551                "metrics".to_owned() => HttpListenerConfig {
552                    addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
553                    authenticator_kind: AuthenticatorKind::None,
554                    enable_tls: false,
555                    routes: HttpRoutesEnabled {
556                        base: RouteGroup::Disabled,
557                        webhook: RouteGroup::Disabled,
558                        internal: RouteGroup::Disabled,
559                        metrics: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
560                        profiling: RouteGroup::Disabled,
561                        mcp_agent: RouteGroup::Disabled,
562                        mcp_developer: RouteGroup::Disabled,
563                        console_config: RouteGroup::Enabled(AllowedRoles::NormalAndInternal),
564                    },
565                },
566            },
567        };
568        self
569    }
570
571    pub fn with_now(mut self, now: NowFn) -> Self {
572        self.now = now;
573        self
574    }
575
576    pub fn with_storage_usage_collection_interval(
577        mut self,
578        storage_usage_collection_interval: Duration,
579    ) -> Self {
580        self.storage_usage_collection_interval = storage_usage_collection_interval;
581        self
582    }
583
584    pub fn with_storage_usage_retention_period(
585        mut self,
586        storage_usage_retention_period: Duration,
587    ) -> Self {
588        self.storage_usage_retention_period = Some(storage_usage_retention_period);
589        self
590    }
591
592    pub fn with_default_cluster_replica_size(
593        mut self,
594        default_cluster_replica_size: String,
595    ) -> Self {
596        self.default_cluster_replica_size = default_cluster_replica_size;
597        self
598    }
599
600    pub fn with_builtin_system_cluster_replica_size(
601        mut self,
602        builtin_system_cluster_replica_size: String,
603    ) -> Self {
604        self.builtin_system_cluster_config.size = builtin_system_cluster_replica_size;
605        self
606    }
607
608    pub fn with_builtin_system_cluster_replication_factor(
609        mut self,
610        builtin_system_cluster_replication_factor: u32,
611    ) -> Self {
612        self.builtin_system_cluster_config.replication_factor =
613            builtin_system_cluster_replication_factor;
614        self
615    }
616
617    pub fn with_builtin_support_cluster_replication_factor(
618        mut self,
619        builtin_support_cluster_replication_factor: u32,
620    ) -> Self {
621        self.builtin_support_cluster_config.replication_factor =
622            builtin_support_cluster_replication_factor;
623        self
624    }
625
626    pub fn with_builtin_catalog_server_cluster_replica_size(
627        mut self,
628        builtin_catalog_server_cluster_replica_size: String,
629    ) -> Self {
630        self.builtin_catalog_server_cluster_config.size =
631            builtin_catalog_server_cluster_replica_size;
632        self
633    }
634
635    pub fn with_propagate_crashes(mut self, propagate_crashes: bool) -> Self {
636        self.propagate_crashes = propagate_crashes;
637        self
638    }
639
640    pub fn with_enable_tracing(mut self, enable_tracing: bool) -> Self {
641        self.enable_tracing = enable_tracing;
642        self
643    }
644
645    pub fn with_bootstrap_role(mut self, bootstrap_role: Option<String>) -> Self {
646        self.bootstrap_role = bootstrap_role;
647        self
648    }
649
650    pub fn with_deploy_generation(mut self, deploy_generation: u64) -> Self {
651        self.deploy_generation = deploy_generation;
652        self
653    }
654
655    pub fn with_system_parameter_default(mut self, param: String, value: String) -> Self {
656        self.system_parameter_defaults.insert(param, value);
657        self
658    }
659
660    pub fn with_mcp_routes(mut self, agent: bool, developer: bool) -> Self {
661        for config in self.listeners_config.http.values_mut() {
662            // Match the MCP routes to the listener's `base` policy (falling back
663            // to `NormalAndInternal` if `base` is disabled, e.g. a metrics-only
664            // listener), then enable/disable them.
665            let roles = config
666                .routes
667                .base
668                .allowed_roles()
669                .unwrap_or(AllowedRoles::NormalAndInternal);
670            let group = |enabled| {
671                if enabled {
672                    RouteGroup::Enabled(roles)
673                } else {
674                    RouteGroup::Disabled
675                }
676            };
677            config.routes.mcp_agent = group(agent);
678            config.routes.mcp_developer = group(developer);
679        }
680        self
681    }
682
683    pub fn with_internal_console_redirect_url(
684        mut self,
685        internal_console_redirect_url: Option<String>,
686    ) -> Self {
687        self.internal_console_redirect_url = internal_console_redirect_url;
688        self
689    }
690
691    pub fn with_metrics_registry(mut self, registry: MetricsRegistry) -> Self {
692        self.metrics_registry = Some(registry);
693        self
694    }
695
696    pub fn with_code_version(mut self, version: semver::Version) -> Self {
697        self.code_version = version;
698        self
699    }
700
701    /// Forces every builtin storage collection through the given migration mechanism,
702    /// `"evolution"` or `"replacement"`.
703    pub fn with_force_builtin_schema_migration(mut self, mechanism: &str) -> Self {
704        self.force_builtin_schema_migration = Some(mechanism.into());
705        self
706    }
707
708    pub fn with_capture(mut self, storage: SharedStorage) -> Self {
709        self.capture = Some(storage);
710        self
711    }
712}
713
714pub struct Listeners {
715    pub inner: crate::Listeners,
716}
717
718impl Listeners {
719    pub async fn new(config: &TestHarness) -> Result<Listeners, anyhow::Error> {
720        let inner = crate::Listeners::bind(config.listeners_config.clone()).await?;
721        Ok(Listeners { inner })
722    }
723
724    pub async fn serve(self, config: TestHarness) -> Result<TestServer, anyhow::Error> {
725        self.serve_with_trigger(config, mz_server_core::cert_reload_never_reload())
726            .await
727    }
728
729    pub async fn serve_with_trigger(
730        self,
731        config: TestHarness,
732        tls_reload_certs: ReloadTrigger,
733    ) -> Result<TestServer, anyhow::Error> {
734        let (data_directory, temp_dir) = match config.data_directory {
735            None => {
736                // If no data directory is provided, we create a temporary
737                // directory. The temporary directory is cleaned up when the
738                // `TempDir` is dropped, so we keep it alive until the `Server` is
739                // dropped.
740                let temp_dir = tempfile::tempdir()?;
741                (temp_dir.path().to_path_buf(), Some(temp_dir))
742            }
743            Some(data_directory) => (data_directory, None),
744        };
745        let scratch_dir = tempfile::tempdir()?;
746        let (consensus_uri, timestamp_oracle_url) = {
747            let seed = config.seed;
748            let cockroach_url = env::var("METADATA_BACKEND_URL")
749                .map_err(|_| anyhow!("METADATA_BACKEND_URL environment variable is not set"))?;
750            let (client, conn) = tokio_postgres::connect(&cockroach_url, NoTls).await?;
751            mz_ore::task::spawn(|| "startup-postgres-conn", async move {
752                if let Err(err) = conn.await {
753                    panic!("connection error: {}", err);
754                };
755            });
756            let consensus_schema = sql!("consensus_{}", seed);
757            let tsoracle_schema = sql!("tsoracle_{}", seed);
758            pg_batch_execute(
759                &client,
760                sql!(
761                    "CREATE SCHEMA IF NOT EXISTS {};
762                     CREATE SCHEMA IF NOT EXISTS {};",
763                    consensus_schema,
764                    tsoracle_schema,
765                ),
766            )
767            .await?;
768            (
769                format!("{cockroach_url}?options=--search_path=consensus_{seed}")
770                    .parse()
771                    .expect("invalid consensus URI"),
772                format!("{cockroach_url}?options=--search_path=tsoracle_{seed}")
773                    .parse()
774                    .expect("invalid timestamp oracle URI"),
775            )
776        };
777        let metrics_registry = config.metrics_registry.unwrap_or_else(MetricsRegistry::new);
778        let orchestrator = ProcessOrchestrator::new(ProcessOrchestratorConfig {
779            image_dir: env::current_exe()?
780                .parent()
781                .unwrap()
782                .parent()
783                .unwrap()
784                .to_path_buf(),
785            suppress_output: false,
786            environment_id: config.environment_id.to_string(),
787            secrets_dir: data_directory.join("secrets"),
788            command_wrapper: vec![],
789            propagate_crashes: config.propagate_crashes,
790            tcp_proxy: None,
791            scratch_directory: scratch_dir.path().to_path_buf(),
792        })
793        .await?;
794        let orchestrator = Arc::new(orchestrator);
795        // Messing with the clock causes persist to expire leases, causing hangs and
796        // panics. Is it possible/desirable to put this back somehow?
797        let persist_now = SYSTEM_TIME.clone();
798        let dyncfgs = mz_dyncfgs::all_dyncfgs();
799
800        let mut updates = ConfigUpdates::default();
801        // Tune down the number of connections to make this all work a little easier
802        // with local postgres.
803        updates.add(&CONSENSUS_CONNECTION_POOL_MAX_SIZE, 1);
804        updates.apply(&dyncfgs);
805
806        let mut persist_cfg = PersistConfig::new(&crate::BUILD_INFO, persist_now.clone(), dyncfgs);
807        persist_cfg.build_version = config.code_version;
808        // Stress persist more by writing rollups frequently
809        persist_cfg.set_rollup_threshold(5);
810
811        let persist_pubsub_server = PersistGrpcPubSubServer::new(&persist_cfg, &metrics_registry);
812        let persist_pubsub_client = persist_pubsub_server.new_same_process_connection();
813        let persist_pubsub_tcp_listener =
814            TcpListener::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0))
815                .await
816                .expect("pubsub addr binding");
817        let persist_pubsub_server_port = persist_pubsub_tcp_listener
818            .local_addr()
819            .expect("pubsub addr has local addr")
820            .port();
821
822        // Spawn the persist pub-sub server.
823        mz_ore::task::spawn(|| "persist_pubsub_server", async move {
824            persist_pubsub_server
825                .serve_with_stream(TcpListenerStream::new(persist_pubsub_tcp_listener))
826                .await
827                .expect("success")
828        });
829        let persist_clients =
830            PersistClientCache::new(persist_cfg, &metrics_registry, |_, _| persist_pubsub_client);
831        let persist_clients = Arc::new(persist_clients);
832        let system_dyncfgs = Arc::clone(&persist_clients.cfg().configs);
833
834        let secrets_controller = Arc::clone(&orchestrator);
835        let mut connection_context = ConnectionContext::for_tests(orchestrator.reader());
836        if !config.aws_connection_context {
837            connection_context.aws_external_id_prefix = None;
838            connection_context.aws_connection_role_arn = None;
839        }
840        let orchestrator = Arc::new(TracingOrchestrator::new(
841            orchestrator,
842            config.orchestrator_tracing_cli_args,
843        ));
844        let tracing_handle = if config.enable_tracing {
845            let config = TracingConfig::<fn(&tracing::Metadata) -> sentry_tracing::EventFilter> {
846                service_name: "environmentd",
847                stderr_log: StderrLogConfig {
848                    format: StderrLogFormat::Json,
849                    filter: EnvFilter::default(),
850                },
851                opentelemetry: Some(OpenTelemetryConfig {
852                    endpoint: "http://fake_address_for_testing:8080".to_string(),
853                    headers: http::HeaderMap::new(),
854                    filter: EnvFilter::default().add_directive(Level::DEBUG.into()),
855                    resource: opentelemetry_sdk::resource::Resource::builder().build(),
856                    max_batch_queue_size: 2048,
857                    max_export_batch_size: 512,
858                    max_concurrent_exports: 1,
859                    batch_scheduled_delay: Duration::from_millis(5000),
860                    max_export_timeout: Duration::from_secs(30),
861                }),
862                tokio_console: None,
863                sentry: None,
864                build_version: crate::BUILD_INFO.version,
865                build_sha: crate::BUILD_INFO.sha,
866                registry: metrics_registry.clone(),
867                capture: config.capture,
868            };
869            mz_ore::tracing::configure(config).await?
870        } else {
871            TracingHandle::disabled()
872        };
873        let host_name = format!(
874            "localhost:{}",
875            self.inner.http["external"].handle.local_addr.port()
876        );
877        let catalog_config = CatalogConfig {
878            persist_clients: Arc::clone(&persist_clients),
879            metrics: Arc::new(mz_catalog::durable::Metrics::new(&MetricsRegistry::new())),
880        };
881
882        let inner = self
883            .inner
884            .serve(crate::Config {
885                catalog_config,
886                timestamp_oracle_url: Some(timestamp_oracle_url),
887                controller: ControllerConfig {
888                    build_info: &crate::BUILD_INFO,
889                    orchestrator,
890                    clusterd_image: "clusterd".into(),
891                    init_container_image: None,
892                    deploy_generation: config.deploy_generation,
893                    persist_location: PersistLocation {
894                        blob_uri: format!("file://{}/persist/blob", data_directory.display())
895                            .parse()
896                            .expect("invalid blob URI"),
897                        consensus_uri,
898                    },
899                    persist_clients,
900                    now: config.now.clone(),
901                    metrics_registry: metrics_registry.clone(),
902                    persist_pubsub_url: format!("http://localhost:{}", persist_pubsub_server_port),
903                    secrets_args: mz_service::secrets::SecretsReaderCliArgs {
904                        secrets_reader: mz_service::secrets::SecretsControllerKind::LocalFile,
905                        secrets_reader_local_file_dir: Some(data_directory.join("secrets")),
906                        secrets_reader_kubernetes_context: None,
907                        secrets_reader_aws_prefix: None,
908                        secrets_reader_name_prefix: None,
909                    },
910                    connection_context,
911                    replica_http_locator: Default::default(),
912                },
913                secrets_controller,
914                cloud_resource_controller: None,
915                system_dyncfgs,
916                tls: config.tls,
917                frontegg: config.frontegg,
918                frontegg_oauth_issuer_url: None,
919                unsafe_mode: config.unsafe_mode,
920                all_features: false,
921                metrics_registry: metrics_registry.clone(),
922                now: config.now,
923                environment_id: config.environment_id,
924                cors_allowed_origin: AllowOrigin::list([]),
925                cors_allowed_origin_list: Vec::new(),
926                cluster_replica_sizes: ClusterReplicaSizeMap::for_tests(),
927                bootstrap_default_cluster_replica_size: config.default_cluster_replica_size,
928                bootstrap_default_cluster_replication_factor: config
929                    .default_cluster_replication_factor,
930                bootstrap_builtin_system_cluster_config: config.builtin_system_cluster_config,
931                bootstrap_builtin_catalog_server_cluster_config: config
932                    .builtin_catalog_server_cluster_config,
933                bootstrap_builtin_probe_cluster_config: config.builtin_probe_cluster_config,
934                bootstrap_builtin_support_cluster_config: config.builtin_support_cluster_config,
935                bootstrap_builtin_analytics_cluster_config: config.builtin_analytics_cluster_config,
936                system_parameter_defaults: config.system_parameter_defaults,
937                availability_zones: Default::default(),
938                tracing_handle,
939                storage_usage_collection_interval: config.storage_usage_collection_interval,
940                storage_usage_retention_period: config.storage_usage_retention_period,
941                segment_api_key: None,
942                segment_client_side: false,
943                test_only_dummy_segment_client: false,
944                egress_addresses: vec![],
945                aws_account_id: None,
946                aws_privatelink_availability_zones: None,
947                launchdarkly_sdk_key: None,
948                launchdarkly_base_uri: None,
949                launchdarkly_key_map: Default::default(),
950                config_sync_file_path: None,
951                config_sync_timeout: Duration::from_secs(30),
952                config_sync_loop_interval: None,
953                bootstrap_role: config.bootstrap_role,
954                http_host_name: Some(host_name),
955                internal_console_redirect_url: config.internal_console_redirect_url,
956                tls_reload_certs,
957                helm_chart_version: None,
958                license_key: ValidatedLicenseKey::for_tests(),
959                external_login_password_mz_system: config.external_login_password_mz_system,
960                force_builtin_schema_migration: config.force_builtin_schema_migration,
961            })
962            .await?;
963
964        Ok(TestServer {
965            inner,
966            metrics_registry,
967            _temp_dir: temp_dir,
968            _scratch_dir: scratch_dir,
969        })
970    }
971}
972
973/// A running instance of `environmentd`.
974pub struct TestServer {
975    pub inner: crate::Server,
976    pub metrics_registry: MetricsRegistry,
977    /// The `TempDir`s are saved to prevent them from being dropped, and thus cleaned up too early.
978    _temp_dir: Option<TempDir>,
979    _scratch_dir: TempDir,
980}
981
982impl TestServer {
983    pub fn connect(&self) -> ConnectBuilder<'_, postgres::NoTls, NoHandle> {
984        ConnectBuilder::new(self).no_tls()
985    }
986
987    pub async fn enable_feature_flags(&self, flags: &[&'static str]) {
988        let internal_client = self.connect().internal().await.unwrap();
989
990        for flag in flags {
991            let query = sql!("ALTER SYSTEM SET {} = true;", Sql::ident(flag));
992            pg_batch_execute(&internal_client, query).await.unwrap();
993        }
994    }
995
996    pub async fn disable_feature_flags(&self, flags: &[&'static str]) {
997        let internal_client = self.connect().internal().await.unwrap();
998
999        for flag in flags {
1000            let query = sql!("ALTER SYSTEM SET {} = false;", Sql::ident(flag));
1001            pg_batch_execute(&internal_client, query).await.unwrap();
1002        }
1003    }
1004
1005    pub fn ws_addr(&self) -> Uri {
1006        format!(
1007            "ws://{}/api/experimental/sql",
1008            self.inner.http_listener_handles["external"].local_addr
1009        )
1010        .parse()
1011        .unwrap()
1012    }
1013
1014    pub fn internal_ws_addr(&self) -> Uri {
1015        format!(
1016            "ws://{}/api/experimental/sql",
1017            self.inner.http_listener_handles["internal"].local_addr
1018        )
1019        .parse()
1020        .unwrap()
1021    }
1022
1023    pub fn http_local_addr(&self) -> SocketAddr {
1024        self.inner.http_listener_handles["external"].local_addr
1025    }
1026
1027    pub fn internal_http_local_addr(&self) -> SocketAddr {
1028        self.inner.http_listener_handles["internal"].local_addr
1029    }
1030
1031    pub fn sql_local_addr(&self) -> SocketAddr {
1032        self.inner.sql_listener_handles["external"].local_addr
1033    }
1034
1035    pub fn internal_sql_local_addr(&self) -> SocketAddr {
1036        self.inner.sql_listener_handles["internal"].local_addr
1037    }
1038}
1039
1040/// A builder struct to configure a pgwire connection to a running [`TestServer`].
1041///
1042/// You can create this struct, and thus open a pgwire connection, using [`TestServer::connect`].
1043pub struct ConnectBuilder<'s, T, H> {
1044    /// A running `environmentd` test server.
1045    server: &'s TestServer,
1046
1047    /// Postgres configuration for connecting to the test server.
1048    pg_config: tokio_postgres::Config,
1049    /// Port to use when connecting to the test server.
1050    port: u16,
1051    /// Tls settings to use.
1052    tls: T,
1053
1054    /// Callback that gets invoked for every notice we receive.
1055    notice_callback: Option<Box<dyn FnMut(tokio_postgres::error::DbError) + Send + 'static>>,
1056
1057    /// Type variable for whether or not we include the handle for the spawned [`tokio::task`].
1058    _with_handle: H,
1059}
1060
1061impl<'s> ConnectBuilder<'s, (), NoHandle> {
1062    fn new(server: &'s TestServer) -> Self {
1063        let mut pg_config = tokio_postgres::Config::new();
1064        pg_config
1065            .host(&Ipv4Addr::LOCALHOST.to_string())
1066            .user("materialize")
1067            .options("--welcome_message=off")
1068            .application_name("environmentd_test_framework");
1069
1070        ConnectBuilder {
1071            server,
1072            pg_config,
1073            port: server.sql_local_addr().port(),
1074            tls: (),
1075            notice_callback: None,
1076            _with_handle: NoHandle,
1077        }
1078    }
1079}
1080
1081impl<'s, T, H> ConnectBuilder<'s, T, H> {
1082    /// Create a pgwire connection without using TLS.
1083    ///
1084    /// Note: this is the default for all connections.
1085    pub fn no_tls(self) -> ConnectBuilder<'s, postgres::NoTls, H> {
1086        ConnectBuilder {
1087            server: self.server,
1088            pg_config: self.pg_config,
1089            port: self.port,
1090            tls: postgres::NoTls,
1091            notice_callback: self.notice_callback,
1092            _with_handle: self._with_handle,
1093        }
1094    }
1095
1096    /// Create a pgwire connection with TLS.
1097    pub fn with_tls<Tls>(self, tls: Tls) -> ConnectBuilder<'s, Tls, H>
1098    where
1099        Tls: MakeTlsConnect<Socket> + Send + 'static,
1100        Tls::TlsConnect: Send,
1101        Tls::Stream: Send,
1102        <Tls::TlsConnect as TlsConnect<Socket>>::Future: Send,
1103    {
1104        ConnectBuilder {
1105            server: self.server,
1106            pg_config: self.pg_config,
1107            port: self.port,
1108            tls,
1109            notice_callback: self.notice_callback,
1110            _with_handle: self._with_handle,
1111        }
1112    }
1113
1114    /// Create a [`ConnectBuilder`] using the provided [`tokio_postgres::Config`].
1115    pub fn with_config(mut self, pg_config: tokio_postgres::Config) -> Self {
1116        self.pg_config = pg_config;
1117        self
1118    }
1119
1120    /// Set the [`SslMode`] to be used with the resulting connection.
1121    pub fn ssl_mode(mut self, mode: SslMode) -> Self {
1122        self.pg_config.ssl_mode(mode);
1123        self
1124    }
1125
1126    /// Set the user for the pgwire connection.
1127    pub fn user(mut self, user: &str) -> Self {
1128        self.pg_config.user(user);
1129        self
1130    }
1131
1132    /// Set the password for the pgwire connection.
1133    pub fn password(mut self, password: &str) -> Self {
1134        self.pg_config.password(password);
1135        self
1136    }
1137
1138    /// Set the application name for the pgwire connection.
1139    pub fn application_name(mut self, application_name: &str) -> Self {
1140        self.pg_config.application_name(application_name);
1141        self
1142    }
1143
1144    /// Set the database name for the pgwire connection.
1145    pub fn dbname(mut self, dbname: &str) -> Self {
1146        self.pg_config.dbname(dbname);
1147        self
1148    }
1149
1150    /// Set the options for the pgwire connection.
1151    pub fn options(mut self, options: &str) -> Self {
1152        self.pg_config.options(options);
1153        self
1154    }
1155
1156    /// Configures this [`ConnectBuilder`] to connect to the __internal__ SQL port of the running
1157    /// [`TestServer`].
1158    ///
1159    /// For example, this will change the port we connect to, and the user we connect as.
1160    pub fn internal(mut self) -> Self {
1161        self.port = self.server.internal_sql_local_addr().port();
1162        self.pg_config.user(mz_sql::session::user::SYSTEM_USER_NAME);
1163        self
1164    }
1165
1166    /// Sets a callback for any database notices that are received from the [`TestServer`].
1167    pub fn notice_callback(self, callback: impl FnMut(DbError) + Send + 'static) -> Self {
1168        ConnectBuilder {
1169            notice_callback: Some(Box::new(callback)),
1170            ..self
1171        }
1172    }
1173
1174    /// Configures this [`ConnectBuilder`] to return the [`mz_ore::task::JoinHandle`] that is
1175    /// polling the underlying postgres connection, associated with the returned client.
1176    pub fn with_handle(self) -> ConnectBuilder<'s, T, WithHandle> {
1177        ConnectBuilder {
1178            server: self.server,
1179            pg_config: self.pg_config,
1180            port: self.port,
1181            tls: self.tls,
1182            notice_callback: self.notice_callback,
1183            _with_handle: WithHandle,
1184        }
1185    }
1186
1187    /// Returns the [`tokio_postgres::Config`] that will be used to connect.
1188    pub fn as_pg_config(&self) -> &tokio_postgres::Config {
1189        &self.pg_config
1190    }
1191}
1192
1193/// This trait enables us to either include or omit the [`mz_ore::task::JoinHandle`] in the result
1194/// of a client connection.
1195pub trait IncludeHandle: Send {
1196    type Output;
1197    fn transform_result(
1198        client: tokio_postgres::Client,
1199        handle: mz_ore::task::JoinHandle<()>,
1200    ) -> Self::Output;
1201}
1202
1203/// Type parameter that denotes we __will not__ return the [`mz_ore::task::JoinHandle`] in the
1204/// result of a [`ConnectBuilder`].
1205pub struct NoHandle;
1206impl IncludeHandle for NoHandle {
1207    type Output = tokio_postgres::Client;
1208    fn transform_result(
1209        client: tokio_postgres::Client,
1210        _handle: mz_ore::task::JoinHandle<()>,
1211    ) -> Self::Output {
1212        client
1213    }
1214}
1215
1216/// Type parameter that denotes we __will__ return the [`mz_ore::task::JoinHandle`] in the result of
1217/// a [`ConnectBuilder`].
1218pub struct WithHandle;
1219impl IncludeHandle for WithHandle {
1220    type Output = (tokio_postgres::Client, mz_ore::task::JoinHandle<()>);
1221    fn transform_result(
1222        client: tokio_postgres::Client,
1223        handle: mz_ore::task::JoinHandle<()>,
1224    ) -> Self::Output {
1225        (client, handle)
1226    }
1227}
1228
1229impl<'s, T, H> IntoFuture for ConnectBuilder<'s, T, H>
1230where
1231    T: MakeTlsConnect<Socket> + Send + 'static,
1232    T::TlsConnect: Send,
1233    T::Stream: Send,
1234    <T::TlsConnect as TlsConnect<Socket>>::Future: Send,
1235    H: IncludeHandle,
1236{
1237    type Output = Result<H::Output, postgres::Error>;
1238    type IntoFuture = BoxFuture<'static, Self::Output>;
1239
1240    fn into_future(mut self) -> Self::IntoFuture {
1241        Box::pin(async move {
1242            assert!(
1243                self.pg_config.get_ports().is_empty(),
1244                "specifying multiple ports is not supported"
1245            );
1246            self.pg_config.port(self.port);
1247
1248            let (client, mut conn) = self.pg_config.connect(self.tls).await?;
1249            let mut notice_callback = self.notice_callback.take();
1250
1251            let handle = task::spawn(|| "connect", async move {
1252                while let Some(msg) = std::future::poll_fn(|cx| conn.poll_message(cx)).await {
1253                    match msg {
1254                        Ok(AsyncMessage::Notice(notice)) => {
1255                            if let Some(callback) = notice_callback.as_mut() {
1256                                callback(notice);
1257                            }
1258                        }
1259                        Ok(msg) => {
1260                            tracing::debug!(?msg, "Dropping message from database");
1261                        }
1262                        Err(e) => {
1263                            // tokio_postgres::Connection docs say:
1264                            // > Return values of None or Some(Err(_)) are “terminal”; callers
1265                            // > should not invoke this method again after receiving one of those
1266                            // > values.
1267                            tracing::info!("connection error: {e}");
1268                            break;
1269                        }
1270                    }
1271                }
1272                tracing::info!("connection closed");
1273            });
1274
1275            let output = H::transform_result(client, handle);
1276            Ok(output)
1277        })
1278    }
1279}
1280
1281/// A running instance of `environmentd`, that exposes blocking/synchronous test helpers.
1282///
1283/// Note: Ideally you should use a [`TestServer`] which relies on an external runtime, e.g. the
1284/// [`tokio::test`] macro. This struct exists so we can incrementally migrate our existing tests.
1285pub struct TestServerWithRuntime {
1286    server: TestServer,
1287    runtime: Arc<Runtime>,
1288}
1289
1290impl TestServerWithRuntime {
1291    /// Returns the [`Runtime`] owned by this [`TestServerWithRuntime`].
1292    ///
1293    /// Can be used to spawn async tasks.
1294    pub fn runtime(&self) -> &Arc<Runtime> {
1295        &self.runtime
1296    }
1297
1298    /// Returns a referece to the inner running `environmentd` [`crate::Server`]`.
1299    pub fn inner(&self) -> &crate::Server {
1300        &self.server.inner
1301    }
1302
1303    /// Connect to the __public__ SQL port of the running `environmentd` server.
1304    pub fn connect<T>(&self, tls: T) -> Result<postgres::Client, postgres::Error>
1305    where
1306        T: MakeTlsConnect<Socket> + Send + 'static,
1307        T::TlsConnect: Send,
1308        T::Stream: Send,
1309        <T::TlsConnect as TlsConnect<Socket>>::Future: Send,
1310    {
1311        self.pg_config().connect(tls)
1312    }
1313
1314    /// Connect to the __internal__ SQL port of the running `environmentd` server.
1315    pub fn connect_internal<T>(&self, tls: T) -> Result<postgres::Client, anyhow::Error>
1316    where
1317        T: MakeTlsConnect<Socket> + Send + 'static,
1318        T::TlsConnect: Send,
1319        T::Stream: Send,
1320        <T::TlsConnect as TlsConnect<Socket>>::Future: Send,
1321    {
1322        Ok(self.pg_config_internal().connect(tls)?)
1323    }
1324
1325    /// Enable LaunchDarkly feature flags.
1326    pub fn enable_feature_flags(&self, flags: &[&'static str]) {
1327        let mut internal_client = self.connect_internal(postgres::NoTls).unwrap();
1328
1329        for flag in flags {
1330            let query = sql!("ALTER SYSTEM SET {} = true;", Sql::ident(flag));
1331            // This uses the synchronous `postgres::Client`; wrappers are async
1332            // and currently only defined for tokio-postgres clients.
1333            #[allow(clippy::disallowed_methods)]
1334            internal_client.batch_execute(query.as_str()).unwrap();
1335        }
1336    }
1337
1338    /// Disable LaunchDarkly feature flags.
1339    pub fn disable_feature_flags(&self, flags: &[&'static str]) {
1340        let mut internal_client = self.connect_internal(postgres::NoTls).unwrap();
1341
1342        for flag in flags {
1343            let query = sql!("ALTER SYSTEM SET {} = false;", Sql::ident(flag));
1344            // This uses the synchronous `postgres::Client`; wrappers are async
1345            // and currently only defined for tokio-postgres clients.
1346            #[allow(clippy::disallowed_methods)]
1347            internal_client.batch_execute(query.as_str()).unwrap();
1348        }
1349    }
1350
1351    /// Return a [`postgres::Config`] for connecting to the __public__ SQL port of the running
1352    /// `environmentd` server.
1353    pub fn pg_config(&self) -> postgres::Config {
1354        let local_addr = self.server.sql_local_addr();
1355        let mut config = postgres::Config::new();
1356        config
1357            .host(&Ipv4Addr::LOCALHOST.to_string())
1358            .port(local_addr.port())
1359            .user("materialize")
1360            .options("--welcome_message=off");
1361        config
1362    }
1363
1364    /// Return a [`postgres::Config`] for connecting to the __internal__ SQL port of the running
1365    /// `environmentd` server.
1366    pub fn pg_config_internal(&self) -> postgres::Config {
1367        let local_addr = self.server.internal_sql_local_addr();
1368        let mut config = postgres::Config::new();
1369        config
1370            .host(&Ipv4Addr::LOCALHOST.to_string())
1371            .port(local_addr.port())
1372            .user("mz_system")
1373            .options("--welcome_message=off");
1374        config
1375    }
1376
1377    pub fn ws_addr(&self) -> Uri {
1378        self.server.ws_addr()
1379    }
1380
1381    pub fn internal_ws_addr(&self) -> Uri {
1382        self.server.internal_ws_addr()
1383    }
1384
1385    pub fn http_local_addr(&self) -> SocketAddr {
1386        self.server.http_local_addr()
1387    }
1388
1389    pub fn internal_http_local_addr(&self) -> SocketAddr {
1390        self.server.internal_http_local_addr()
1391    }
1392
1393    pub fn sql_local_addr(&self) -> SocketAddr {
1394        self.server.sql_local_addr()
1395    }
1396
1397    pub fn internal_sql_local_addr(&self) -> SocketAddr {
1398        self.server.internal_sql_local_addr()
1399    }
1400
1401    /// Returns the metrics registry for the test server.
1402    pub fn metrics_registry(&self) -> &MetricsRegistry {
1403        &self.server.metrics_registry
1404    }
1405}
1406
1407#[derive(Debug, Clone, Copy, Eq, PartialEq, Ord, PartialOrd)]
1408pub struct MzTimestamp(pub u64);
1409
1410impl<'a> FromSql<'a> for MzTimestamp {
1411    fn from_sql(ty: &Type, raw: &'a [u8]) -> Result<MzTimestamp, Box<dyn Error + Sync + Send>> {
1412        let n = mz_pgrepr::Numeric::from_sql(ty, raw)?;
1413        Ok(MzTimestamp(u64::try_from(n.0.0)?))
1414    }
1415
1416    fn accepts(ty: &Type) -> bool {
1417        mz_pgrepr::Numeric::accepts(ty)
1418    }
1419}
1420
1421pub trait PostgresErrorExt {
1422    fn unwrap_db_error(self) -> DbError;
1423}
1424
1425impl PostgresErrorExt for postgres::Error {
1426    fn unwrap_db_error(self) -> DbError {
1427        match self.source().and_then(|e| e.downcast_ref::<DbError>()) {
1428            Some(e) => e.clone(),
1429            None => panic!("expected DbError, but got: {:?}", self),
1430        }
1431    }
1432}
1433
1434impl<T, E> PostgresErrorExt for Result<T, E>
1435where
1436    E: PostgresErrorExt,
1437{
1438    fn unwrap_db_error(self) -> DbError {
1439        match self {
1440            Ok(_) => panic!("expected Err(DbError), but got Ok(_)"),
1441            Err(e) => e.unwrap_db_error(),
1442        }
1443    }
1444}
1445
1446/// Group commit will block writes until the current time has advanced. This can make
1447/// performing inserts while using deterministic time difficult. This is a helper
1448/// method to perform writes and advance the current time.
1449pub async fn insert_with_deterministic_timestamps(
1450    table: &'static str,
1451    values: &'static str,
1452    server: &TestServer,
1453    now: Arc<std::sync::Mutex<EpochMillis>>,
1454) -> Result<(), Box<dyn Error>> {
1455    let client_write = server.connect().await?;
1456    let client_read = server.connect().await?;
1457
1458    let mut current_ts = get_explain_timestamp(table, &client_read).await;
1459
1460    let insert_query = format!("INSERT INTO {} VALUES {values}", Sql::ident(table));
1461
1462    // The `values` fragment is raw SQL text in test code and cannot currently
1463    // be represented as a composable `Sql` fragment.
1464    #[allow(clippy::disallowed_methods)]
1465    let write_future = client_write.execute(&insert_query, &[]);
1466    let timestamp_interval = tokio::time::interval(Duration::from_millis(1));
1467
1468    let mut write_future = std::pin::pin!(write_future);
1469    let mut timestamp_interval = std::pin::pin!(timestamp_interval);
1470
1471    // Keep increasing `now` until the write has executed succeed. Table advancements may
1472    // have increased the global timestamp by an unknown amount.
1473    loop {
1474        tokio::select! {
1475            _ = (&mut write_future) => return Ok(()),
1476            _ = timestamp_interval.tick() => {
1477                current_ts += 1;
1478                *now.lock().expect("lock poisoned") = current_ts;
1479            }
1480        };
1481    }
1482}
1483
1484pub async fn get_explain_timestamp(from_suffix: &str, client: &Client) -> EpochMillis {
1485    try_get_explain_timestamp(from_suffix, client)
1486        .await
1487        .unwrap()
1488}
1489
1490pub async fn try_get_explain_timestamp(
1491    from_suffix: &str,
1492    client: &Client,
1493) -> Result<EpochMillis, anyhow::Error> {
1494    let det = get_explain_timestamp_determination(from_suffix, client).await?;
1495    let ts = det.determination.timestamp_context.timestamp_or_default();
1496    Ok(ts.into())
1497}
1498
1499pub async fn get_explain_timestamp_determination(
1500    from_suffix: &str,
1501    client: &Client,
1502) -> Result<TimestampExplanation, anyhow::Error> {
1503    // `from_suffix` is a raw SQL suffix used by this test helper and cannot
1504    // currently be represented as a composable `Sql` fragment.
1505    #[allow(clippy::disallowed_methods)]
1506    let row = client
1507        .query_one(
1508            &format!("EXPLAIN TIMESTAMP AS JSON FOR SELECT * FROM {from_suffix}"),
1509            &[],
1510        )
1511        .await?;
1512    let explain: String = row.get(0);
1513    Ok(serde_json::from_str(&explain).unwrap())
1514}
1515
1516/// Helper function to create a Postgres source.
1517///
1518/// IMPORTANT: Make sure to call closure that is returned at the end of the test to clean up
1519/// Postgres state.
1520///
1521/// WARNING: If multiple tests use this, and the tests are run in parallel, then make sure the test
1522/// use different postgres tables.
1523pub async fn create_postgres_source_with_table<'a>(
1524    server: &TestServer,
1525    mz_client: &Client,
1526    table_name: &str,
1527    table_schema: &str,
1528    source_name: &str,
1529) -> (
1530    Client,
1531    impl FnOnce(&'a Client, &'a Client) -> LocalBoxFuture<'a, ()>,
1532) {
1533    server
1534        .enable_feature_flags(&["enable_create_table_from_source"])
1535        .await;
1536
1537    let postgres_url = env::var("POSTGRES_URL")
1538        .map_err(|_| anyhow!("POSTGRES_URL environment variable is not set"))
1539        .unwrap();
1540
1541    let (pg_client, connection) = tokio_postgres::connect(&postgres_url, postgres::NoTls)
1542        .await
1543        .unwrap();
1544
1545    let pg_config: tokio_postgres::Config = postgres_url.parse().unwrap();
1546    let user = pg_config.get_user().unwrap_or("postgres");
1547    let db_name = pg_config.get_dbname().unwrap_or(user);
1548    let ports = pg_config.get_ports();
1549    let port = if ports.is_empty() { 5432 } else { ports[0] };
1550    let hosts = pg_config.get_hosts();
1551    let host = if hosts.is_empty() {
1552        "localhost".to_string()
1553    } else {
1554        match &hosts[0] {
1555            Host::Tcp(host) => host.to_string(),
1556            Host::Unix(host) => host.to_str().unwrap().to_string(),
1557        }
1558    };
1559    let password = pg_config.get_password();
1560
1561    mz_ore::task::spawn(|| "postgres-source-connection", async move {
1562        if let Err(e) = connection.await {
1563            panic!("connection error: {}", e);
1564        }
1565    });
1566
1567    // Create table in Postgres with publication.
1568    let _ = pg_execute(
1569        &pg_client,
1570        sql!("DROP TABLE IF EXISTS {};", Sql::ident(table_name)),
1571        &[],
1572    )
1573    .await
1574    .unwrap();
1575    let _ = pg_execute(
1576        &pg_client,
1577        sql!("DROP PUBLICATION IF EXISTS {};", Sql::ident(source_name)),
1578        &[],
1579    )
1580    .await
1581    .unwrap();
1582    // `table_schema` is a raw schema fragment in this test helper and cannot
1583    // currently be represented as a composable `Sql` fragment.
1584    #[allow(clippy::disallowed_methods)]
1585    let _ = pg_client
1586        .execute(
1587            format!("CREATE TABLE {} {table_schema};", Sql::ident(table_name)).as_str(),
1588            &[],
1589        )
1590        .await
1591        .unwrap();
1592    let _ = pg_execute(
1593        &pg_client,
1594        sql!(
1595            "ALTER TABLE {} REPLICA IDENTITY FULL;",
1596            Sql::ident(table_name)
1597        ),
1598        &[],
1599    )
1600    .await
1601    .unwrap();
1602    let _ = pg_execute(
1603        &pg_client,
1604        sql!(
1605            "CREATE PUBLICATION {} FOR TABLE {};",
1606            Sql::ident(source_name),
1607            Sql::ident(table_name)
1608        ),
1609        &[],
1610    )
1611    .await
1612    .unwrap();
1613
1614    // Create postgres source in Materialize.
1615    let mut connection_str = format!("HOST '{host}', PORT {port}, USER {user}, DATABASE {db_name}");
1616    if let Some(password) = password {
1617        let password = std::str::from_utf8(password).unwrap();
1618        pg_batch_execute(
1619            mz_client,
1620            sql!("CREATE SECRET s AS {}", Sql::literal(password)),
1621        )
1622        .await
1623        .unwrap();
1624        connection_str = format!("{connection_str}, PASSWORD SECRET s");
1625    }
1626    // `connection_str` is a raw connection-option fragment generated for tests
1627    // and cannot currently be represented as a composable `Sql` fragment.
1628    #[allow(clippy::disallowed_methods)]
1629    mz_client
1630        .batch_execute(format!("CREATE CONNECTION pgconn TO POSTGRES ({connection_str})").as_str())
1631        .await
1632        .unwrap();
1633    pg_batch_execute(
1634        mz_client,
1635        sql!(
1636            "CREATE SOURCE {} \
1637             FROM POSTGRES \
1638             CONNECTION pgconn \
1639             (PUBLICATION {})",
1640            Sql::ident(source_name),
1641            Sql::literal(source_name),
1642        ),
1643    )
1644    .await
1645    .unwrap();
1646    pg_batch_execute(
1647        mz_client,
1648        sql!(
1649            "CREATE TABLE {} \
1650             FROM SOURCE {} \
1651             (REFERENCE {});",
1652            Sql::ident(table_name),
1653            Sql::ident(source_name),
1654            Sql::ident(table_name),
1655        ),
1656    )
1657    .await
1658    .unwrap();
1659
1660    let table_name = table_name.to_string();
1661    let source_name = source_name.to_string();
1662    (
1663        pg_client,
1664        move |mz_client: &'a Client, pg_client: &'a Client| {
1665            let f: Pin<Box<dyn Future<Output = ()> + 'a>> = Box::pin(async move {
1666                pg_batch_execute(
1667                    mz_client,
1668                    sql!("DROP SOURCE {} CASCADE;", Sql::ident(&source_name)),
1669                )
1670                .await
1671                .unwrap();
1672                pg_batch_execute(mz_client, sql!("DROP CONNECTION pgconn;"))
1673                    .await
1674                    .unwrap();
1675
1676                let _ = pg_execute(
1677                    pg_client,
1678                    sql!("DROP PUBLICATION {};", Sql::ident(&source_name)),
1679                    &[],
1680                )
1681                .await
1682                .unwrap();
1683                let _ = pg_execute(
1684                    pg_client,
1685                    sql!("DROP TABLE {};", Sql::ident(&table_name)),
1686                    &[],
1687                )
1688                .await
1689                .unwrap();
1690            });
1691            f
1692        },
1693    )
1694}
1695
1696pub async fn wait_for_pg_table_population(mz_client: &Client, view_name: &str, source_rows: i64) {
1697    let current_isolation = pg_query_one(mz_client, sql!("SHOW transaction_isolation"), &[])
1698        .await
1699        .unwrap()
1700        .get::<_, String>(0);
1701    pg_batch_execute(mz_client, sql!("SET transaction_isolation = SERIALIZABLE"))
1702        .await
1703        .unwrap();
1704    Retry::default()
1705        .retry_async(|_| async move {
1706            let rows = pg_query_one(
1707                mz_client,
1708                sql!("SELECT COUNT(*) FROM {};", Sql::ident(view_name)),
1709                &[],
1710            )
1711            .await
1712            .unwrap()
1713            .get::<_, i64>(0);
1714            if rows == source_rows {
1715                Ok(())
1716            } else {
1717                Err(format!(
1718                    "Waiting for {source_rows} row to be ingested. Currently at {rows}."
1719                ))
1720            }
1721        })
1722        .await
1723        .unwrap();
1724    pg_batch_execute(
1725        mz_client,
1726        sql!(
1727            "SET transaction_isolation = {}",
1728            Sql::literal(&current_isolation),
1729        ),
1730    )
1731    .await
1732    .unwrap();
1733}
1734
1735// Initializes a websocket connection. Returns the init messages before the initial ReadyForQuery.
1736pub fn auth_with_ws(
1737    ws: &mut WebSocket<MaybeTlsStream<TcpStream>>,
1738    mut options: BTreeMap<String, String>,
1739) -> Result<Vec<WebSocketResponse>, anyhow::Error> {
1740    if !options.contains_key("welcome_message") {
1741        options.insert("welcome_message".into(), "off".into());
1742    }
1743    auth_with_ws_impl(
1744        ws,
1745        Message::Text(
1746            serde_json::to_string(&WebSocketAuth::Basic {
1747                user: "materialize".into(),
1748                password: "".into(),
1749                options,
1750            })
1751            .unwrap()
1752            .into(),
1753        ),
1754    )
1755}
1756
1757pub fn auth_with_ws_impl(
1758    ws: &mut WebSocket<MaybeTlsStream<TcpStream>>,
1759    auth_message: Message,
1760) -> Result<Vec<WebSocketResponse>, anyhow::Error> {
1761    ws.send(auth_message)?;
1762
1763    // Wait for initial ready response.
1764    let mut msgs = Vec::new();
1765    loop {
1766        let resp = ws.read()?;
1767        match resp {
1768            Message::Text(msg) => {
1769                let msg: WebSocketResponse = serde_json::from_str(&msg).unwrap();
1770                match msg {
1771                    WebSocketResponse::ReadyForQuery(_) => break,
1772                    msg => {
1773                        msgs.push(msg);
1774                    }
1775                }
1776            }
1777            Message::Ping(_) => continue,
1778            Message::Close(None) => return Err(anyhow!("ws closed after auth")),
1779            Message::Close(Some(close_frame)) => {
1780                return Err(anyhow!("ws closed after auth").context(close_frame));
1781            }
1782            _ => panic!("unexpected response: {:?}", resp),
1783        }
1784    }
1785    Ok(msgs)
1786}
1787
1788pub fn make_header<H: Header>(h: H) -> HeaderMap {
1789    let mut map = HeaderMap::new();
1790    map.typed_insert(h);
1791    map
1792}
1793
1794pub fn make_pg_tls<F>(configure: F) -> MakeTlsConnector
1795where
1796    F: FnOnce(&mut SslConnectorBuilder) -> Result<(), ErrorStack>,
1797{
1798    let mut connector_builder = SslConnector::builder(SslMethod::tls()).unwrap();
1799    // Disable TLS v1.3 because `postgres` and `hyper` produce stabler error
1800    // messages with TLS v1.2.
1801    //
1802    // Briefly, in TLS v1.3, failing to present a client certificate does not
1803    // error during the TLS handshake, as it does in TLS v1.2, but on the first
1804    // attempt to read from the stream. But both `postgres` and `hyper` write a
1805    // bunch of data before attempting to read from the stream. With a failed
1806    // TLS v1.3 connection, sometimes `postgres` and `hyper` succeed in writing
1807    // out this data, and then return a nice error message on the call to read.
1808    // But sometimes the connection is closed before they write out the data,
1809    // and so they report "connection closed" before they ever call read, never
1810    // noticing the underlying SSL error.
1811    //
1812    // It's unclear who's bug this is. Is it on `hyper`/`postgres` to call read
1813    // if writing to the stream fails to see if a TLS error occured? Is it on
1814    // OpenSSL to provide a better API [1]? Is it a protocol issue that ought to
1815    // be corrected in TLS v1.4? We don't want to answer these questions, so we
1816    // just avoid TLS v1.3 for now.
1817    //
1818    // [1]: https://github.com/openssl/openssl/issues/11118
1819    let options = connector_builder.options() | SslOptions::NO_TLSV1_3;
1820    connector_builder.set_options(options);
1821    configure(&mut connector_builder).unwrap();
1822    MakeTlsConnector::new(connector_builder.build())
1823}
1824
1825/// A certificate authority for use in tests.
1826pub struct Ca {
1827    pub dir: TempDir,
1828    pub name: X509Name,
1829    pub cert: X509,
1830    pub pkey: PKey<Private>,
1831}
1832
1833impl Ca {
1834    fn make_ca(name: &str, parent: Option<&Ca>) -> Result<Ca, Box<dyn Error>> {
1835        let dir = tempfile::tempdir()?;
1836        let rsa = Rsa::generate(2048)?;
1837        let pkey = PKey::from_rsa(rsa)?;
1838        let name = {
1839            let mut builder = X509NameBuilder::new()?;
1840            builder.append_entry_by_nid(Nid::COMMONNAME, name)?;
1841            builder.build()
1842        };
1843        let cert = {
1844            let mut builder = X509::builder()?;
1845            builder.set_version(2)?;
1846            builder.set_pubkey(&pkey)?;
1847            builder.set_issuer_name(parent.map(|ca| &ca.name).unwrap_or(&name))?;
1848            builder.set_subject_name(&name)?;
1849            builder.set_not_before(&*Asn1Time::days_from_now(0)?)?;
1850            builder.set_not_after(&*Asn1Time::days_from_now(365)?)?;
1851            builder.append_extension(BasicConstraints::new().critical().ca().build()?)?;
1852            builder.sign(
1853                parent.map(|ca| &ca.pkey).unwrap_or(&pkey),
1854                MessageDigest::sha256(),
1855            )?;
1856            builder.build()
1857        };
1858        fs::write(dir.path().join("ca.crt"), cert.to_pem()?)?;
1859        Ok(Ca {
1860            dir,
1861            name,
1862            cert,
1863            pkey,
1864        })
1865    }
1866
1867    /// Creates a new root certificate authority.
1868    pub fn new_root(name: &str) -> Result<Ca, Box<dyn Error>> {
1869        Ca::make_ca(name, None)
1870    }
1871
1872    /// Returns the path to the CA's certificate.
1873    pub fn ca_cert_path(&self) -> PathBuf {
1874        self.dir.path().join("ca.crt")
1875    }
1876
1877    /// Requests a new intermediate certificate authority.
1878    pub fn request_ca(&self, name: &str) -> Result<Ca, Box<dyn Error>> {
1879        Ca::make_ca(name, Some(self))
1880    }
1881
1882    /// Generates a certificate with the specified Common Name (CN) that is
1883    /// signed by the CA.
1884    ///
1885    /// Returns the paths to the certificate and key.
1886    pub fn request_client_cert(&self, name: &str) -> Result<(PathBuf, PathBuf), Box<dyn Error>> {
1887        self.request_cert(name, iter::empty())
1888    }
1889
1890    /// Like `request_client_cert`, but permits specifying additional IP
1891    /// addresses to attach as Subject Alternate Names.
1892    pub fn request_cert<I>(&self, name: &str, ips: I) -> Result<(PathBuf, PathBuf), Box<dyn Error>>
1893    where
1894        I: IntoIterator<Item = IpAddr>,
1895    {
1896        let rsa = Rsa::generate(2048)?;
1897        let pkey = PKey::from_rsa(rsa)?;
1898        let subject_name = {
1899            let mut builder = X509NameBuilder::new()?;
1900            builder.append_entry_by_nid(Nid::COMMONNAME, name)?;
1901            builder.build()
1902        };
1903        let cert = {
1904            let mut builder = X509::builder()?;
1905            builder.set_version(2)?;
1906            builder.set_pubkey(&pkey)?;
1907            builder.set_issuer_name(self.cert.subject_name())?;
1908            builder.set_subject_name(&subject_name)?;
1909            builder.set_not_before(&*Asn1Time::days_from_now(0)?)?;
1910            builder.set_not_after(&*Asn1Time::days_from_now(365)?)?;
1911            for ip in ips {
1912                builder.append_extension(
1913                    SubjectAlternativeName::new()
1914                        .ip(&ip.to_string())
1915                        .build(&builder.x509v3_context(None, None))?,
1916                )?;
1917            }
1918            builder.sign(&self.pkey, MessageDigest::sha256())?;
1919            builder.build()
1920        };
1921        let cert_path = self.dir.path().join(Path::new(name).with_extension("crt"));
1922        let key_path = self.dir.path().join(Path::new(name).with_extension("key"));
1923        fs::write(&cert_path, cert.to_pem()?)?;
1924        fs::write(&key_path, pkey.private_key_to_pem_pkcs8()?)?;
1925        Ok((cert_path, key_path))
1926    }
1927}
1928
1929/// Sums the counter series named `name` whose labels include all of `labels`.
1930///
1931/// Returns 0 when no series matches, which is how a labelled counter reads
1932/// before its first increment.
1933pub fn get_counter_value(registry: &MetricsRegistry, name: &str, labels: &[(&str, &str)]) -> u64 {
1934    let Some(family) = registry.gather().into_iter().find(|m| m.name() == name) else {
1935        return 0;
1936    };
1937    family
1938        .get_metric()
1939        .iter()
1940        .filter(|metric| {
1941            labels.iter().all(|(name, value)| {
1942                metric
1943                    .get_label()
1944                    .iter()
1945                    .any(|label| label.name() == *name && label.value() == *value)
1946            })
1947        })
1948        .map(|metric| u64::cast_lossy(metric.get_counter().value()))
1949        .sum()
1950}