1use 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#[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 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 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 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 pub async fn start(self) -> TestServer {
262 self.try_start().await.expect("Failed to start test Server")
263 }
264
265 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 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 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 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 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 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 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 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 let persist_now = SYSTEM_TIME.clone();
798 let dyncfgs = mz_dyncfgs::all_dyncfgs();
799
800 let mut updates = ConfigUpdates::default();
801 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 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 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
973pub struct TestServer {
975 pub inner: crate::Server,
976 pub metrics_registry: MetricsRegistry,
977 _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
1040pub struct ConnectBuilder<'s, T, H> {
1044 server: &'s TestServer,
1046
1047 pg_config: tokio_postgres::Config,
1049 port: u16,
1051 tls: T,
1053
1054 notice_callback: Option<Box<dyn FnMut(tokio_postgres::error::DbError) + Send + 'static>>,
1056
1057 _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 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 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 pub fn with_config(mut self, pg_config: tokio_postgres::Config) -> Self {
1116 self.pg_config = pg_config;
1117 self
1118 }
1119
1120 pub fn ssl_mode(mut self, mode: SslMode) -> Self {
1122 self.pg_config.ssl_mode(mode);
1123 self
1124 }
1125
1126 pub fn user(mut self, user: &str) -> Self {
1128 self.pg_config.user(user);
1129 self
1130 }
1131
1132 pub fn password(mut self, password: &str) -> Self {
1134 self.pg_config.password(password);
1135 self
1136 }
1137
1138 pub fn application_name(mut self, application_name: &str) -> Self {
1140 self.pg_config.application_name(application_name);
1141 self
1142 }
1143
1144 pub fn dbname(mut self, dbname: &str) -> Self {
1146 self.pg_config.dbname(dbname);
1147 self
1148 }
1149
1150 pub fn options(mut self, options: &str) -> Self {
1152 self.pg_config.options(options);
1153 self
1154 }
1155
1156 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 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 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 pub fn as_pg_config(&self) -> &tokio_postgres::Config {
1189 &self.pg_config
1190 }
1191}
1192
1193pub 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
1203pub 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
1216pub 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 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
1281pub struct TestServerWithRuntime {
1286 server: TestServer,
1287 runtime: Arc<Runtime>,
1288}
1289
1290impl TestServerWithRuntime {
1291 pub fn runtime(&self) -> &Arc<Runtime> {
1295 &self.runtime
1296 }
1297
1298 pub fn inner(&self) -> &crate::Server {
1300 &self.server.inner
1301 }
1302
1303 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 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 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 #[allow(clippy::disallowed_methods)]
1334 internal_client.batch_execute(query.as_str()).unwrap();
1335 }
1336 }
1337
1338 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 #[allow(clippy::disallowed_methods)]
1347 internal_client.batch_execute(query.as_str()).unwrap();
1348 }
1349 }
1350
1351 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 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 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
1446pub 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 #[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 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 #[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
1516pub 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 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 #[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 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 #[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(¤t_isolation),
1729 ),
1730 )
1731 .await
1732 .unwrap();
1733}
1734
1735pub 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 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 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
1825pub 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 pub fn new_root(name: &str) -> Result<Ca, Box<dyn Error>> {
1869 Ca::make_ca(name, None)
1870 }
1871
1872 pub fn ca_cert_path(&self) -> PathBuf {
1874 self.dir.path().join("ca.crt")
1875 }
1876
1877 pub fn request_ca(&self, name: &str) -> Result<Ca, Box<dyn Error>> {
1879 Ca::make_ca(name, Some(self))
1880 }
1881
1882 pub fn request_client_cert(&self, name: &str) -> Result<(PathBuf, PathBuf), Box<dyn Error>> {
1887 self.request_cert(name, iter::empty())
1888 }
1889
1890 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
1929pub 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}