1use crate::{MetadataExecutor, PgMetaError, bootstrap};
2use ahri_tre_types::DataStoreProfile;
3use native_tls::{Certificate, TlsConnector};
4use postgres::{Client, Transaction};
5use postgres_native_tls::MakeTlsConnector;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
8enum TlsVerification {
9 None,
10 Certificate,
11 CertificateAndHostname,
12}
13
14pub const M0001_INITIAL_SCHEMA_VERSION: &str = "m0001_initial_schema";
16
17#[derive(Debug, Clone, PartialEq, Eq)]
19pub struct DirectPostgresConnectConfig {
20 pub host: String,
21 pub port: u16,
22 pub dbname: String,
23 pub username: String,
24 pub password: Option<String>,
25 pub sslmode: Option<String>,
26 pub connect_timeout_secs: Option<u64>,
27}
28
29impl DirectPostgresConnectConfig {
30 pub fn from_profile(
32 profile: &DataStoreProfile,
33 username: impl Into<String>,
34 password: Option<String>,
35 ) -> Self {
36 Self {
37 host: profile.server.clone(),
38 port: profile.port,
39 dbname: profile.dbname.clone(),
40 username: username.into(),
41 password,
42 sslmode: profile.sslmode.clone(),
43 connect_timeout_secs: Some(10),
44 }
45 }
46
47 pub fn conninfo(&self) -> String {
49 self.conninfo_with_sslmode(self.sslmode.as_deref())
50 }
51
52 fn postgres_conninfo(&self) -> String {
53 self.postgres_conninfo_with_routing_address(None)
54 }
55
56 fn postgres_conninfo_with_routing_address(&self, routing_address: Option<&str>) -> String {
57 let sslmode = self.sslmode.as_deref().map(|sslmode| {
61 if matches!(
62 sslmode.to_ascii_lowercase().as_str(),
63 "verify-ca" | "verify-full"
64 ) {
65 "require"
66 } else {
67 sslmode
68 }
69 });
70 self.conninfo_with_sslmode_and_routing_address(sslmode, routing_address)
71 }
72
73 fn conninfo_with_sslmode(&self, sslmode: Option<&str>) -> String {
74 self.conninfo_with_sslmode_and_routing_address(sslmode, None)
75 }
76
77 fn conninfo_with_sslmode_and_routing_address(
78 &self,
79 sslmode: Option<&str>,
80 routing_address: Option<&str>,
81 ) -> String {
82 let mut parts = vec![conninfo_pair("host", &self.host)];
83 if let Some(routing_address) = routing_address.filter(|value| !value.is_empty()) {
84 parts.push(conninfo_pair("hostaddr", routing_address));
85 }
86 parts.extend([
87 conninfo_pair("port", &self.port.to_string()),
88 conninfo_pair("dbname", &self.dbname),
89 conninfo_pair("user", &self.username),
90 ]);
91
92 if let Some(sslmode) = sslmode.filter(|value| !value.is_empty()) {
93 parts.push(conninfo_pair("sslmode", sslmode));
94 }
95
96 if let Some(timeout) = self.connect_timeout_secs {
97 parts.push(conninfo_pair("connect_timeout", &timeout.to_string()));
98 }
99
100 if let Some(password) = self.password.as_deref().filter(|value| !value.is_empty()) {
101 parts.push(conninfo_pair("password", password));
102 }
103
104 parts.join(" ")
105 }
106
107 pub fn safe_conninfo(&self) -> String {
109 let mut redacted = self.clone();
110 if redacted.password.is_some() {
111 redacted.password = Some("<redacted>".to_string());
112 }
113 redacted.conninfo()
114 }
115}
116
117#[derive(Debug, Clone)]
119pub struct PgMetadataAdapter {
120 pub connection_info: String,
121 pub connection_description: String,
122 root_certificates: Vec<String>,
123 tls_verification: TlsVerification,
124}
125
126impl PgMetadataAdapter {
127 pub fn new(connection_info: impl Into<String>) -> Self {
129 let connection_info = connection_info.into();
130 Self {
131 connection_description: connection_info.clone(),
132 connection_info,
133 root_certificates: Vec::new(),
134 tls_verification: TlsVerification::None,
135 }
136 }
137
138 pub fn direct(config: DirectPostgresConnectConfig) -> Self {
140 let tls_verification = tls_verification(config.sslmode.as_deref());
141 Self {
142 connection_info: config.postgres_conninfo(),
143 connection_description: config.safe_conninfo(),
144 root_certificates: Vec::new(),
145 tls_verification,
146 }
147 }
148
149 pub fn direct_with_ca(
151 config: DirectPostgresConnectConfig,
152 root_certificates: Vec<String>,
153 ) -> Self {
154 let tls_verification = tls_verification(config.sslmode.as_deref());
155 Self {
156 connection_info: config.postgres_conninfo(),
157 connection_description: config.safe_conninfo(),
158 root_certificates,
159 tls_verification,
160 }
161 }
162
163 pub fn direct_with_ca_and_routing_address(
166 config: DirectPostgresConnectConfig,
167 routing_address: Option<&str>,
168 root_certificates: Vec<String>,
169 ) -> Self {
170 let tls_verification = tls_verification(config.sslmode.as_deref());
171 let connection_info = config.postgres_conninfo_with_routing_address(routing_address);
172 let mut redacted = config.clone();
173 if redacted.password.is_some() {
174 redacted.password = Some("<redacted>".to_string());
175 }
176 Self {
177 connection_info,
178 connection_description: redacted.conninfo_with_sslmode_and_routing_address(
179 config.sslmode.as_deref(),
180 routing_address,
181 ),
182 root_certificates,
183 tls_verification,
184 }
185 }
186
187 pub fn connection_description(&self) -> &str {
189 &self.connection_description
190 }
191
192 pub fn connect_observed(
195 &self,
196 context: Option<&ahri_tre_observability::CorrelationContext>,
197 ) -> Result<PgMetadataConnection, PgMetaError> {
198 use ahri_tre_observability::{FailureCategory, Outcome, Stage};
199 let span = context.map(|context| context.span(Stage::MetadataConnection));
200 let result = self.connect();
201 if let Some(span) = span {
202 let category = result.as_ref().err().map(PgMetaError::operational_category);
203 span.finish(
204 match category {
205 None => Outcome::Success,
206 Some(FailureCategory::Authentication | FailureCategory::Authorization) => {
207 Outcome::Rejected
208 }
209 Some(_) => Outcome::Unavailable,
210 },
211 category,
212 );
213 }
214 result.map(|mut connection| {
215 connection.read_observation = context.map(|context| MetadataReadObservation {
216 span: Some(context.span(Stage::MetadataQuery)),
217 failure: std::cell::Cell::new(None),
218 });
219 connection
220 })
221 }
222
223 pub fn connect(&self) -> Result<PgMetadataConnection, PgMetaError> {
230 let tls = tls_connector(self.tls_verification, &self.root_certificates)?;
231 let client = Client::connect(&self.connection_info, MakeTlsConnector::new(tls))
232 .map_err(PgMetaError::Connect)?;
233 Ok(PgMetadataConnection {
234 client,
235 admission_active: false,
236 read_observation: None,
237 connection_description: self.connection_description.clone(),
238 })
239 }
240}
241
242pub struct PgMetadataConnection {
244 pub(crate) client: Client,
245 pub(crate) admission_active: bool,
246 read_observation: Option<MetadataReadObservation>,
247 connection_description: String,
248}
249
250impl std::fmt::Debug for PgMetadataConnection {
251 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
252 f.debug_struct("PgMetadataConnection")
253 .field("connection_description", &self.connection_description)
254 .finish_non_exhaustive()
255 }
256}
257
258impl PgMetadataConnection {
259 pub fn observe_error(&self, error: &PgMetaError) {
262 if let Some(observation) = &self.read_observation
263 && observation.failure.get().is_none()
264 {
265 observation.failure.set(Some(error.operational_category()));
266 }
267 }
268
269 pub fn connection_description(&self) -> &str {
271 &self.connection_description
272 }
273
274 pub fn client(&mut self) -> &mut Client {
276 &mut self.client
277 }
278
279 pub fn health_check(&mut self) -> Result<PgHealth, PgMetaError> {
285 let row = MetadataExecutor::query_one(
286 self,
287 "metadata health check",
288 "SELECT current_database(), current_user, current_setting('server_version')",
289 &[],
290 )?;
291
292 let migration_row = MetadataExecutor::query_optional(
293 self,
294 "metadata schema bootstrap check",
295 "SELECT 1
296 FROM information_schema.tables
297 WHERE table_schema = 'public'
298 AND table_name = 'ahri_tre_schema_migrations'",
299 &[],
300 )?;
301
302 Ok(PgHealth {
303 database: row.get(0),
304 user: row.get(1),
305 server_version: row.get(2),
306 schema_bootstrapped: migration_row.is_some(),
307 schema_compatibility: crate::schema::datastore_schema_status(self)
308 .map(|status| status.status.as_str().to_string())
309 .ok(),
310 })
311 }
312
313 pub fn bootstrap_schema(&mut self) -> Result<SchemaBootstrapReport, PgMetaError> {
320 let shape = self
321 .client
322 .query_one(
323 "SELECT
324 to_regclass('public.ahri_tre_schema_migrations') IS NOT NULL,
325 EXISTS (
326 SELECT 1
327 FROM information_schema.tables
328 WHERE table_schema = 'public'
329 AND table_name <> 'ahri_tre_schema_migrations'
330 )",
331 &[],
332 )
333 .map_err(PgMetaError::Query)?;
334 let history_present: bool = shape.get(0);
335 let other_tables_present: bool = shape.get(1);
336 if history_present {
337 let versions = self
338 .client
339 .query(
340 "SELECT version::text FROM ahri_tre_schema_migrations ORDER BY version",
341 &[],
342 )
343 .map_err(PgMetaError::Query)?
344 .into_iter()
345 .map(|row| row.get::<_, String>(0))
346 .collect::<Vec<_>>();
347 if versions.as_slice()
348 == [
349 M0001_INITIAL_SCHEMA_VERSION,
350 crate::CURRENT_DATASTORE_SCHEMA_VERSION,
351 ]
352 {
353 return Ok(SchemaBootstrapReport {
354 version: crate::CURRENT_DATASTORE_SCHEMA_VERSION.into(),
355 applied: false,
356 });
357 }
358 if versions.as_slice() != [M0001_INITIAL_SCHEMA_VERSION] {
359 return Err(unsupported_bootstrap_shape());
360 }
361 } else if other_tables_present {
362 return Err(unsupported_bootstrap_shape());
363 }
364
365 let mut transaction = self
366 .client
367 .transaction()
368 .map_err(PgMetaError::Transaction)?;
369
370 transaction
371 .batch_execute(
372 "
373 SET search_path TO public;
374 CREATE TABLE IF NOT EXISTS ahri_tre_schema_migrations (
375 version TEXT PRIMARY KEY,
376 description TEXT NOT NULL,
377 applied_at TIMESTAMPTZ NOT NULL DEFAULT now(),
378 scope TEXT NOT NULL DEFAULT 'metadata'
379 );
380 ALTER TABLE ahri_tre_schema_migrations
381 ADD COLUMN IF NOT EXISTS scope TEXT NOT NULL DEFAULT 'metadata';
382 ",
383 )
384 .map_err(PgMetaError::Query)?;
385
386 bootstrap::bootstrap_metadata_schema(&mut transaction)?;
387
388 let baseline_inserted = transaction
389 .execute(
390 "
391 INSERT INTO ahri_tre_schema_migrations (version, description, scope)
392 VALUES ($1, $2, 'metadata')
393 ON CONFLICT (version) DO NOTHING
394 ",
395 &[
396 &M0001_INITIAL_SCHEMA_VERSION,
397 &"Complete AHRI TRE version-1 metadata schema",
398 ],
399 )
400 .map_err(PgMetaError::Query)?;
401
402 transaction.commit().map_err(PgMetaError::Transaction)?;
403
404 Ok(SchemaBootstrapReport {
405 version: M0001_INITIAL_SCHEMA_VERSION.to_string(),
406 applied: baseline_inserted == 1,
407 })
408 }
409
410 pub fn with_metadata_transaction<T, F>(&mut self, f: F) -> Result<T, PgMetaError>
417 where
418 F: FnOnce(&mut Transaction<'_>) -> Result<T, PgMetaError>,
419 {
420 let mut transaction = self
421 .client
422 .transaction()
423 .map_err(PgMetaError::Transaction)?;
424 let result = f(&mut transaction)?;
425 transaction.commit().map_err(PgMetaError::Transaction)?;
426 Ok(result)
427 }
428}
429
430fn unsupported_bootstrap_shape() -> PgMetaError {
431 PgMetaError::Decode {
432 field: "schema_version",
433 value: "<unsupported baseline>".to_string(),
434 message: "existing database is not the clean version-1 baseline".to_string(),
435 }
436}
437
438#[derive(Debug, Clone, PartialEq, Eq)]
440pub struct PgHealth {
441 pub database: String,
442 pub user: String,
443 pub server_version: String,
444 pub schema_bootstrapped: bool,
445 pub schema_compatibility: Option<String>,
446}
447
448#[derive(Debug, Clone, PartialEq, Eq)]
450pub struct SchemaBootstrapReport {
451 pub version: String,
452 pub applied: bool,
453}
454
455fn conninfo_pair(key: &str, value: &str) -> String {
456 format!("{key}={}", quote_conninfo_value(value))
457}
458
459fn quote_conninfo_value(value: &str) -> String {
460 let mut escaped = String::with_capacity(value.len() + 2);
461 escaped.push('\'');
462
463 for ch in value.chars() {
464 if ch == '\\' || ch == '\'' {
465 escaped.push('\\');
466 }
467 escaped.push(ch);
468 }
469
470 escaped.push('\'');
471 escaped
472}
473
474fn tls_connector(
475 verification: TlsVerification,
476 root_certificates: &[String],
477) -> Result<TlsConnector, PgMetaError> {
478 let mut builder = TlsConnector::builder();
479 match verification {
480 TlsVerification::None => {
481 builder.danger_accept_invalid_certs(true);
482 builder.danger_accept_invalid_hostnames(true);
483 }
484 TlsVerification::Certificate => {
485 builder.danger_accept_invalid_hostnames(true);
486 }
487 TlsVerification::CertificateAndHostname => {}
488 }
489 for certificate in root_certificates {
490 builder.add_root_certificate(
491 Certificate::from_pem(certificate.as_bytes()).map_err(PgMetaError::Tls)?,
492 );
493 }
494
495 builder.build().map_err(PgMetaError::Tls)
496}
497
498fn tls_verification(sslmode: Option<&str>) -> TlsVerification {
499 match sslmode.map(str::to_ascii_lowercase).as_deref() {
500 Some("verify-ca") => TlsVerification::Certificate,
501 Some("verify-full") => TlsVerification::CertificateAndHostname,
502 _ => TlsVerification::None,
503 }
504}
505
506struct MetadataReadObservation {
507 span: Option<ahri_tre_observability::OperationSpan>,
508 failure: std::cell::Cell<Option<ahri_tre_observability::FailureCategory>>,
509}
510
511impl Drop for MetadataReadObservation {
512 fn drop(&mut self) {
513 use ahri_tre_observability::{FailureCategory, Outcome};
514 if let Some(span) = self.span.take() {
515 let category = self.failure.get();
516 let outcome = match category {
517 None if std::thread::panicking() => Outcome::InternalFailure,
518 None => Outcome::Success,
519 Some(FailureCategory::Authentication | FailureCategory::Authorization) => {
520 Outcome::Rejected
521 }
522 Some(_) => Outcome::Unavailable,
523 };
524 span.finish(outcome, category);
525 }
526 }
527}
528
529#[cfg(test)]
530mod tests {
531 use super::{DirectPostgresConnectConfig, PgMetadataAdapter, TlsVerification};
532 use std::str::FromStr;
533
534 #[test]
535 fn verify_full_uses_supported_postgres_mode_and_strict_tls() {
536 let config = DirectPostgresConnectConfig {
537 host: "postgres".to_string(),
538 port: 5432,
539 dbname: "tre".to_string(),
540 username: "tre_fixture".to_string(),
541 password: Some("fixture-password".to_string()),
542 sslmode: Some("verify-full".to_string()),
543 connect_timeout_secs: Some(10),
544 };
545
546 let adapter = PgMetadataAdapter::direct_with_ca(config, vec!["fixture CA".to_string()]);
547
548 assert!(adapter.connection_info.contains("sslmode='require'"));
549 assert!(!adapter.connection_info.contains("verify-full"));
550 assert!(
551 adapter
552 .connection_description
553 .contains("sslmode='verify-full'")
554 );
555 assert_eq!(
556 adapter.tls_verification,
557 TlsVerification::CertificateAndHostname
558 );
559 assert_eq!(adapter.root_certificates, vec!["fixture CA".to_string()]);
560 postgres::Config::from_str(&adapter.connection_info)
561 .expect("rust-postgres should accept the translated connection settings");
562 }
563
564 #[test]
565 fn verify_ca_disables_only_hostname_verification() {
566 assert_eq!(
567 super::tls_verification(Some("verify-ca")),
568 TlsVerification::Certificate
569 );
570 assert_eq!(
571 super::tls_verification(Some("VERIFY-FULL")),
572 TlsVerification::CertificateAndHostname
573 );
574 assert_eq!(
575 super::tls_verification(Some("require")),
576 TlsVerification::None
577 );
578 }
579
580 #[test]
581 fn routing_address_is_used_for_transport_without_replacing_tls_host() {
582 let config = DirectPostgresConnectConfig {
583 host: "postgres.example.test".to_string(),
584 port: 5432,
585 dbname: "tre".to_string(),
586 username: "tre_fixture".to_string(),
587 password: Some("fixture-password".to_string()),
588 sslmode: Some("verify-full".to_string()),
589 connect_timeout_secs: Some(10),
590 };
591
592 let adapter = PgMetadataAdapter::direct_with_ca_and_routing_address(
593 config,
594 Some("10.0.0.12"),
595 vec!["fixture CA".to_string()],
596 );
597
598 assert!(
599 adapter
600 .connection_info
601 .contains("host='postgres.example.test'")
602 );
603 assert!(adapter.connection_info.contains("hostaddr='10.0.0.12'"));
604 assert!(
605 adapter
606 .connection_description
607 .contains("hostaddr='10.0.0.12'")
608 );
609 assert!(!adapter.connection_description.contains("fixture-password"));
610 postgres::Config::from_str(&adapter.connection_info)
611 .expect("rust-postgres accepts separate host and transport routing");
612 }
613
614 #[test]
615 fn direct_config_builds_quoted_conninfo_and_redacts_password() {
616 let config = DirectPostgresConnectConfig {
617 host: "tre-postgres".to_string(),
618 port: 5432,
619 dbname: "tre data".to_string(),
620 username: "tre_user".to_string(),
621 password: Some("pa'ss\\word".to_string()),
622 sslmode: Some("disable".to_string()),
623 connect_timeout_secs: Some(5),
624 };
625
626 let conninfo = config.conninfo();
627 assert!(conninfo.contains("dbname='tre data'"));
628 assert!(conninfo.contains("password='pa\\'ss\\\\word'"));
629 assert!(conninfo.contains("connect_timeout='5'"));
630
631 let safe_conninfo = config.safe_conninfo();
632 assert!(!safe_conninfo.contains("pa\\'ss\\\\word"));
633 assert!(safe_conninfo.contains("password='<redacted>'"));
634 }
635
636 #[test]
637 fn adapter_keeps_raw_and_safe_connection_descriptions_separate() {
638 let config = DirectPostgresConnectConfig {
639 host: "tre-postgres".to_string(),
640 port: 5432,
641 dbname: "tre".to_string(),
642 username: "tre_user".to_string(),
643 password: Some("secret".to_string()),
644 sslmode: None,
645 connect_timeout_secs: None,
646 };
647
648 let adapter = PgMetadataAdapter::direct(config);
649
650 assert!(adapter.connection_info.contains("password='secret'"));
651 assert!(!adapter.connection_description.contains("secret"));
652 assert_eq!(
653 adapter.connection_description(),
654 "host='tre-postgres' port='5432' dbname='tre' user='tre_user' password='<redacted>'"
655 );
656 }
657}