diff --git a/Cargo.lock b/Cargo.lock index a9b989f87..a029b4d53 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -10928,6 +10928,7 @@ dependencies = [ "thiserror 2.0.19", "tokio", "tokio-test", + "tower 0.5.3", "tracing", "urlencoding", "utoipa", @@ -10986,6 +10987,7 @@ dependencies = [ "sea-orm", "serde", "serde_json", + "serial_test", "sha2 0.11.0", "tar", "temps-auth", @@ -10999,10 +11001,12 @@ dependencies = [ "time", "tokio", "tokio-util", + "tower 0.5.3", "tracing", "utoipa", "uuid", "webpki-roots 1.0.9", + "wiremock", "x509-parser", ] diff --git a/crates/temps-audit/src/plugin.rs b/crates/temps-audit/src/plugin.rs index 79d660500..b588106c8 100644 --- a/crates/temps-audit/src/plugin.rs +++ b/crates/temps-audit/src/plugin.rs @@ -5,7 +5,7 @@ use std::sync::Arc; use temps_core::plugin::{ PluginContext, PluginError, PluginRoutes, ServiceRegistrationContext, TempsPlugin, }; -use temps_core::AuditLogger; +use temps_core::{AuditLogger, AuditLoggerSlot}; use utoipa::OpenApi; use crate::{handlers, AuditService}; @@ -42,7 +42,10 @@ impl TempsPlugin for AuditPlugin { // Create AuditService let audit_service = Arc::new(AuditService::new(db.clone(), ip_address_service.clone())); context.register_service(audit_service.clone()); - let audit_trait: Arc = audit_service.clone(); + let initial_logger: Arc = audit_service.clone(); + let audit_slot = Arc::new(AuditLoggerSlot::new(initial_logger)); + context.register_service(audit_slot.clone()); + let audit_trait: Arc = audit_slot; context.register_service(audit_trait); tracing::debug!("Audit plugin services registered successfully"); diff --git a/crates/temps-auth/src/permissions.rs b/crates/temps-auth/src/permissions.rs index 058a94861..c7c7f1cf2 100644 --- a/crates/temps-auth/src/permissions.rs +++ b/crates/temps-auth/src/permissions.rs @@ -82,6 +82,11 @@ pub enum Permission { SettingsRead, SettingsWrite, + // DNS provider and unattended automation permissions + DnsProvidersRead, + DnsProvidersWrite, + DnsAutomationWrite, + // Files permissions FilesRead, FilesWrite, @@ -323,6 +328,9 @@ impl fmt::Display for Permission { Permission::WebSocketProxyConnect => "websocket_proxy:connect", Permission::SettingsRead => "settings:read", Permission::SettingsWrite => "settings:write", + Permission::DnsProvidersRead => "dns_providers:read", + Permission::DnsProvidersWrite => "dns_providers:write", + Permission::DnsAutomationWrite => "dns_automation:write", Permission::ErrorTrackingRead => "error_tracking:read", Permission::ErrorTrackingWrite => "error_tracking:write", Permission::ErrorTrackingCreate => "error_tracking:create", @@ -437,6 +445,9 @@ impl Permission { "external_services:create" => Some(Permission::ExternalServicesCreate), "settings:read" => Some(Permission::SettingsRead), "settings:write" => Some(Permission::SettingsWrite), + "dns_providers:read" => Some(Permission::DnsProvidersRead), + "dns_providers:write" => Some(Permission::DnsProvidersWrite), + "dns_automation:write" => Some(Permission::DnsAutomationWrite), "files:read" => Some(Permission::FilesRead), "files:write" => Some(Permission::FilesWrite), "files:delete" => Some(Permission::FilesDelete), @@ -583,6 +594,9 @@ impl Permission { Permission::ExternalServicesCreate, Permission::SettingsRead, Permission::SettingsWrite, + Permission::DnsProvidersRead, + Permission::DnsProvidersWrite, + Permission::DnsAutomationWrite, Permission::FilesRead, Permission::FilesWrite, Permission::FilesDelete, @@ -807,6 +821,9 @@ impl Role { Permission::SessionMetricsRead, Permission::SettingsRead, Permission::SettingsWrite, + Permission::DnsProvidersRead, + Permission::DnsProvidersWrite, + Permission::DnsAutomationWrite, Permission::SecretsRead, Permission::SpeedInsightsRead, Permission::SystemAdmin, @@ -946,6 +963,9 @@ impl Role { Permission::SessionMetricsRead, Permission::SettingsRead, Permission::SettingsWrite, + Permission::DnsProvidersRead, + Permission::DnsProvidersWrite, + Permission::DnsAutomationWrite, Permission::SpeedInsightsRead, Permission::SystemAdmin, Permission::SystemRead, @@ -1306,6 +1326,23 @@ mod tests { assert!(admin_permissions.contains(&Permission::EmailsSend)); } + #[test] + fn dns_governance_permissions_round_trip_and_stay_admin_only() { + for (permission, serialized) in [ + (Permission::DnsProvidersRead, "dns_providers:read"), + (Permission::DnsProvidersWrite, "dns_providers:write"), + (Permission::DnsAutomationWrite, "dns_automation:write"), + ] { + assert_eq!(permission.to_string(), serialized); + assert_eq!(Permission::from_str(serialized), Some(permission)); + assert!(Permission::all().contains(&permission)); + assert!(Role::Admin.has_permission(&permission)); + assert!(Role::PlatformAdmin.has_permission(&permission)); + assert!(!Role::User.has_permission(&permission)); + assert!(!Role::Reader.has_permission(&permission)); + } + } + #[test] fn test_user_has_email_permissions() { let user_permissions = Role::User.permissions(); diff --git a/crates/temps-cli/src/commands/serve/console.rs b/crates/temps-cli/src/commands/serve/console.rs index 07df2f3f9..ccdc9cc04 100644 --- a/crates/temps-cli/src/commands/serve/console.rs +++ b/crates/temps-cli/src/commands/serve/console.rs @@ -1842,6 +1842,11 @@ pub async fn start_console_api(params: ConsoleApiParams) -> anyhow::Result<()> { service_context.register_service(encryption_service.clone()); service_context.register_service(cookie_crypto.clone()); service_context.register_service(docker.clone()); + // Background DNS mutation is fail-closed until an optional policy plugin + // claims this slot. DomainsPlugin captures the slot before later plugins + // register, so the indirection must exist before plugin initialization. + let dns_automation_gate_slot = Arc::new(temps_core::DnsAutomationGateSlot::new()); + service_context.register_service(dns_automation_gate_slot); // Pre-registered before any plugin runs so ProxyPlugin uses this exact // slot instance instead of creating its own — see the field doc on // `ConsoleApiParams::retention_resolver_slot`. @@ -1898,10 +1903,11 @@ pub async fn start_console_api(params: ConsoleApiParams) -> anyhow::Result<()> { // (depends only on ServerConfig for the data dir). Registered early so // every later plugin can require the Arc. debug!("Registering TelemetryPlugin"); - let telemetry_plugin = Box::new(TelemetryPlugin::new( - config.clone(), - env!("CARGO_PKG_VERSION"), - )); + // TEMPS_VERSION (git-describe, set by build.rs) is used instead of + // CARGO_PKG_VERSION so nightly/beta builds report a version telemetry + // can actually distinguish from a tagged release -- CARGO_PKG_VERSION + // is the static Cargo.toml version and is identical across all of them. + let telemetry_plugin = Box::new(TelemetryPlugin::new(config.clone(), env!("TEMPS_VERSION"))); plugin_manager.register_plugin(telemetry_plugin); // 2. QueuePlugin - registers the pre-created job queue into the service context diff --git a/crates/temps-core/src/audit.rs b/crates/temps-core/src/audit.rs index 09fb0f808..041ee70b3 100644 --- a/crates/temps-core/src/audit.rs +++ b/crates/temps-core/src/audit.rs @@ -1,5 +1,7 @@ use anyhow::Result; use serde::Serialize; +use std::sync::{Arc, RwLock}; +use thiserror::Error; /// Context information common to all audit events #[derive(Debug, Clone, Serialize)] @@ -45,3 +47,136 @@ pub trait AuditLogger: Send + Sync { /// Creates an audit log entry for the given operation async fn create_audit_log(&self, operation: &dyn AuditOperation) -> Result<()>; } + +/// Stable indirection for audit consumers constructed before optional +/// decorators register. Replacing the target updates every previously-captured +/// `Arc`, so decorators registered later cannot be bypassed. +pub struct AuditLoggerSlot { + target: RwLock>, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)] +pub enum AuditLoggerSlotError { + #[error("audit logger slot read lock is poisoned")] + ReadLockPoisoned, + #[error("audit logger slot write lock is poisoned")] + WriteLockPoisoned, +} + +impl AuditLoggerSlot { + pub fn new(target: Arc) -> Self { + Self { + target: RwLock::new(target), + } + } + + pub fn current(&self) -> std::result::Result, AuditLoggerSlotError> { + self.target + .read() + .map(|target| target.clone()) + .map_err(|_| AuditLoggerSlotError::ReadLockPoisoned) + } + + pub fn replace( + &self, + target: Arc, + ) -> std::result::Result<(), AuditLoggerSlotError> { + let mut current = self + .target + .write() + .map_err(|_| AuditLoggerSlotError::WriteLockPoisoned)?; + *current = target; + Ok(()) + } +} + +#[async_trait::async_trait] +impl AuditLogger for AuditLoggerSlot { + async fn create_audit_log(&self, operation: &dyn AuditOperation) -> Result<()> { + let target = self.current().map_err(anyhow::Error::new)?; + target.create_audit_log(operation).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicUsize, Ordering}; + + struct TestOperation; + + impl AuditOperation for TestOperation { + fn operation_type(&self) -> String { + "TEST".to_string() + } + fn user_id(&self) -> Option { + None + } + fn ip_address(&self) -> Option { + None + } + fn user_agent(&self) -> &str { + "test" + } + fn serialize(&self) -> Result { + Ok("{}".to_string()) + } + } + + struct CountingLogger(Arc); + + #[async_trait::async_trait] + impl AuditLogger for CountingLogger { + async fn create_audit_log(&self, _operation: &dyn AuditOperation) -> Result<()> { + self.0.fetch_add(1, Ordering::SeqCst); + Ok(()) + } + } + + #[tokio::test] + async fn captured_trait_object_routes_to_replacement() { + let original_count = Arc::new(AtomicUsize::new(0)); + let replacement_count = Arc::new(AtomicUsize::new(0)); + let original: Arc = Arc::new(CountingLogger(original_count.clone())); + let slot = Arc::new(AuditLoggerSlot::new(original)); + let captured: Arc = slot.clone(); + + captured.create_audit_log(&TestOperation).await.unwrap(); + slot.replace(Arc::new(CountingLogger(replacement_count.clone()))) + .unwrap(); + captured.create_audit_log(&TestOperation).await.unwrap(); + + assert_eq!(original_count.load(Ordering::SeqCst), 1); + assert_eq!(replacement_count.load(Ordering::SeqCst), 1); + } + + #[test] + fn poisoned_slot_returns_typed_errors() { + fn poisoned_slot() -> Arc { + let logger: Arc = + Arc::new(CountingLogger(Arc::new(AtomicUsize::new(0)))); + let slot = Arc::new(AuditLoggerSlot::new(logger)); + let worker_slot = slot.clone(); + let _ = std::thread::spawn(move || { + let _write_guard = worker_slot.target.write().unwrap(); + panic!("poison audit logger slot for test"); + }) + .join(); + slot + } + + let read_slot = poisoned_slot(); + assert!(matches!( + read_slot.current(), + Err(AuditLoggerSlotError::ReadLockPoisoned) + )); + + let write_slot = poisoned_slot(); + let replacement: Arc = + Arc::new(CountingLogger(Arc::new(AtomicUsize::new(0)))); + assert!(matches!( + write_slot.replace(replacement), + Err(AuditLoggerSlotError::WriteLockPoisoned) + )); + } +} diff --git a/crates/temps-core/src/dns_automation.rs b/crates/temps-core/src/dns_automation.rs new file mode 100644 index 000000000..6d2b6126c --- /dev/null +++ b/crates/temps-core/src/dns_automation.rs @@ -0,0 +1,204 @@ +//! Authorization seam for unattended DNS mutations. +//! +//! DNS provider credentials can modify production infrastructure. Human API +//! access is protected separately by `temps-auth` permissions; this module +//! governs background work, where no authenticated principal exists. The +//! default is deliberately fail-closed: a plugin must explicitly install a +//! gate before the certificate scheduler may publish ACME DNS-01 records. + +use std::sync::{Arc, OnceLock}; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use thiserror::Error; + +/// The only background DNS purpose currently supported by Temps. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum DnsAutomationPurpose { + AcmeDns01, +} + +/// Exact DNS replacement authorized for an unattended operation. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct DnsAutomationMutation { + pub record_type: String, + pub name: String, + pub value: String, +} + +/// Context for an unattended ACME DNS-01 mutation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct DnsAutomationRequest { + pub purpose: DnsAutomationPurpose, + pub domain: String, + pub zone: String, + pub provider_id: i32, + pub provider_name: String, + /// Each entry means replace stale values at this exact name, then publish + /// the supplied value. No broader zone mutation is authorized. + pub mutations: Vec, +} + +/// Result of evaluating an unattended DNS mutation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum DnsAutomationDecision { + Allow, + Deny { reason: String }, +} + +/// Typed failure returned when an automation policy cannot be evaluated. +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub enum DnsAutomationError { + #[error( + "DNS automation policy evaluation failed for provider {provider_id} ({provider_name}), domain {domain}, zone {zone}: {reason}" + )] + PolicyEvaluationFailed { + provider_id: i32, + provider_name: String, + domain: String, + zone: String, + reason: String, + }, +} + +impl DnsAutomationError { + pub fn policy_evaluation_failed( + request: &DnsAutomationRequest, + reason: impl Into, + ) -> Self { + Self::PolicyEvaluationFailed { + provider_id: request.provider_id, + provider_name: request.provider_name.clone(), + domain: request.domain.clone(), + zone: request.zone.clone(), + reason: reason.into(), + } + } +} + +/// Policy extension point for unattended DNS mutations. +#[async_trait] +pub trait DnsAutomationGate: Send + Sync { + async fn authorize( + &self, + request: &DnsAutomationRequest, + ) -> Result; +} + +/// Deferred, write-once gate used across the plugin registration boundary. +/// +/// Core services are constructed before optional plugins register. They keep +/// this slot, while an authorized plugin may claim it later. An unclaimed slot +/// denies automation, so missing or failed plugin registration cannot expose +/// DNS credentials to background mutation. +#[derive(Default)] +pub struct DnsAutomationGateSlot { + gate: OnceLock>, +} + +impl DnsAutomationGateSlot { + pub fn new() -> Self { + Self::default() + } + + /// Install the process-wide DNS automation policy. Only the first caller + /// succeeds; later callers receive their gate back unchanged. + pub fn set(&self, gate: Arc) -> Result<(), Arc> { + self.gate.set(gate) + } +} + +#[async_trait] +impl DnsAutomationGate for DnsAutomationGateSlot { + async fn authorize( + &self, + request: &DnsAutomationRequest, + ) -> Result { + let Some(gate) = self.gate.get() else { + return Ok(DnsAutomationDecision::Deny { + reason: "unattended DNS automation is not enabled for this installation" + .to_string(), + }); + }; + + gate.authorize(request).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + + struct AllowGate; + + #[async_trait] + impl DnsAutomationGate for AllowGate { + async fn authorize( + &self, + _request: &DnsAutomationRequest, + ) -> Result { + Ok(DnsAutomationDecision::Allow) + } + } + + fn request() -> DnsAutomationRequest { + DnsAutomationRequest { + purpose: DnsAutomationPurpose::AcmeDns01, + domain: "example.com".to_string(), + zone: "example.com".to_string(), + provider_id: 7, + provider_name: "production-dns".to_string(), + mutations: vec![DnsAutomationMutation { + record_type: "TXT".to_string(), + name: "_acme-challenge.example.com".to_string(), + value: "challenge-token".to_string(), + }], + } + } + + #[tokio::test] + async fn unclaimed_slot_denies_automation() { + let slot = DnsAutomationGateSlot::new(); + let decision = slot.authorize(&request()).await.unwrap(); + + assert!(matches!(decision, DnsAutomationDecision::Deny { .. })); + } + + #[tokio::test] + async fn claimed_slot_delegates_to_registered_gate() { + let slot = DnsAutomationGateSlot::new(); + assert!(slot.set(Arc::new(AllowGate)).is_ok()); + + assert_eq!( + slot.authorize(&request()).await.unwrap(), + DnsAutomationDecision::Allow + ); + } + + #[test] + fn slot_cannot_be_replaced_after_it_is_claimed() { + let slot = DnsAutomationGateSlot::new(); + assert!(slot.set(Arc::new(AllowGate)).is_ok()); + assert!(slot.set(Arc::new(AllowGate)).is_err()); + } + + #[test] + fn policy_error_carries_provider_and_zone_context() { + let request = request(); + let error = DnsAutomationError::policy_evaluation_failed(&request, "policy store offline"); + + assert!(matches!( + error, + DnsAutomationError::PolicyEvaluationFailed { + provider_id: 7, + ref provider_name, + ref domain, + ref zone, + ref reason, + } if provider_name == "production-dns" + && domain == "example.com" + && zone == "example.com" + && reason == "policy store offline" + )); + } +} diff --git a/crates/temps-core/src/lib.rs b/crates/temps-core/src/lib.rs index 8d880902e..781db0af9 100644 --- a/crates/temps-core/src/lib.rs +++ b/crates/temps-core/src/lib.rs @@ -5,6 +5,7 @@ pub mod audit; pub mod client_ip; pub mod config; pub mod deployment; +pub mod dns_automation; pub mod env_vars_provider; pub mod error; pub mod error_builder; @@ -55,6 +56,10 @@ pub use client_ip::resolve_client_ip; pub use config::*; pub use constants::*; pub use deployment::*; +pub use dns_automation::{ + DnsAutomationDecision, DnsAutomationError, DnsAutomationGate, DnsAutomationGateSlot, + DnsAutomationMutation, DnsAutomationPurpose, DnsAutomationRequest, +}; pub use env_vars_provider::{ flatten_integration_env_vars, IntegrationEnvVar, IntegrationServiceInfo, ProjectEnvVarsProvider, ProjectIntegrationEnvVars, diff --git a/crates/temps-dns/Cargo.toml b/crates/temps-dns/Cargo.toml index ac1b82bae..90ed7e750 100644 --- a/crates/temps-dns/Cargo.toml +++ b/crates/temps-dns/Cargo.toml @@ -73,6 +73,7 @@ temps-migrations = { path = "../temps-migrations" } wiremock = "0.6" serial_test = "4.0" mockall = "0.15" +tower.workspace = true # ADR-024 integration test: in-process DNS client to query the real # control-plane resolver over a UDP socket (mirrors temps-dns-resolver's # own end_to_end.rs client). diff --git a/crates/temps-dns/src/errors.rs b/crates/temps-dns/src/errors.rs index fb57e621d..6f2edac3d 100644 --- a/crates/temps-dns/src/errors.rs +++ b/crates/temps-dns/src/errors.rs @@ -8,6 +8,12 @@ pub enum DnsError { #[error("Provider not found: {0}")] ProviderNotFound(i32), + #[error("DNS provider {provider_id} ({provider_name}) is inactive")] + ProviderInactive { + provider_id: i32, + provider_name: String, + }, + #[error("Invalid provider type: {0}")] InvalidProviderType(String), @@ -26,6 +32,26 @@ pub enum DnsError { #[error("Domain not found: {0}")] DomainNotFound(String), + #[error( + "Managed DNS domain '{requested_domain}' canonicalizes to '{canonical_domain}', which is already managed by domain ID {existing_managed_domain_id} on provider {existing_provider_id}" + )] + ManagedDomainAlreadyExists { + requested_domain: String, + canonical_domain: String, + existing_managed_domain_id: i32, + existing_provider_id: i32, + }, + + #[error( + "Ambiguous managed DNS zone '{canonical_zone}' for requested domain '{requested_domain}': managed domain IDs {managed_domain_ids:?} use provider IDs {provider_ids:?}" + )] + AmbiguousManagedDomain { + requested_domain: String, + canonical_zone: String, + managed_domain_ids: Vec, + provider_ids: Vec, + }, + #[error("Record not found: {0}")] RecordNotFound(String), diff --git a/crates/temps-dns/src/handlers/mod.rs b/crates/temps-dns/src/handlers/mod.rs index 9ef42e38b..3aad3213b 100644 --- a/crates/temps-dns/src/handlers/mod.rs +++ b/crates/temps-dns/src/handlers/mod.rs @@ -40,14 +40,15 @@ use crate::services::{ /// Audit record for managed-domain write operations. #[derive(Debug, Clone, serde::Serialize)] -struct ManagedDomainAudit { +struct DnsGovernanceAudit { context: AuditContext, provider_id: i32, domain: String, action: String, + details: serde_json::Value, } -impl AuditOperation for ManagedDomainAudit { +impl AuditOperation for DnsGovernanceAudit { fn operation_type(&self) -> String { self.action.clone() } @@ -300,6 +301,10 @@ fn default_true() -> bool { true } +fn managed_domain_automation_enabled(auto_manage: bool, sync_generated_records: bool) -> bool { + auto_manage || sync_generated_records +} + /// Request to update a managed domain's settings. #[derive(Debug, Clone, Deserialize, ToSchema)] pub struct UpdateManagedDomainApiRequest { @@ -427,9 +432,26 @@ impl From for Problem { DnsError::ProviderNotFound(id) => problemdetails::new(StatusCode::NOT_FOUND) .with_title("Provider Not Found") .with_detail(format!("DNS provider with ID {} not found", id)), + DnsError::ProviderInactive { + provider_id, + provider_name, + } => problemdetails::new(StatusCode::BAD_REQUEST) + .with_title("DNS Provider Is Inactive") + .with_detail(format!( + "DNS provider {} ({}) is inactive and cannot perform this operation", + provider_id, provider_name + )), DnsError::DomainNotFound(domain) => problemdetails::new(StatusCode::NOT_FOUND) .with_title("Domain Not Found") .with_detail(format!("Domain {} not found", domain)), + DnsError::ManagedDomainAlreadyExists { .. } => { + problemdetails::new(StatusCode::CONFLICT) + .with_title("Managed DNS Domain Already Exists") + .with_detail(error.to_string()) + } + DnsError::AmbiguousManagedDomain { .. } => problemdetails::new(StatusCode::CONFLICT) + .with_title("Ambiguous Managed DNS Zone") + .with_detail(error.to_string()), DnsError::ZoneNotFound(zone) => problemdetails::new(StatusCode::NOT_FOUND) .with_title("Zone Not Found") .with_detail(format!("DNS zone {} not found", zone)), @@ -484,7 +506,7 @@ async fn list_dns_providers( RequireAuth(auth): RequireAuth, State(state): State>, ) -> Result { - permission_check!(auth, Permission::SettingsRead); + permission_check!(auth, Permission::DnsProvidersRead); let providers = state.provider_service.list().await?; @@ -536,9 +558,10 @@ async fn list_dns_providers( async fn create_dns_provider( RequireAuth(auth): RequireAuth, State(state): State>, + Extension(metadata): Extension, Json(request): Json, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); let credentials: ProviderCredentials = request.credentials.into(); @@ -567,8 +590,8 @@ async fn create_dns_provider( let response = DnsProviderResponse { id: provider.id, - name: provider.name, - provider_type: provider.provider_type, + name: provider.name.clone(), + provider_type: provider.provider_type.clone(), credentials: masked_creds, is_active: provider.is_active, description: provider.description, @@ -579,6 +602,20 @@ async fn create_dns_provider( updated_at: provider.updated_at.to_rfc3339(), }; + log_dns_governance_audit( + &state, + &auth, + &metadata, + provider.id, + "", + "DNS_PROVIDER_CREATED", + serde_json::json!({ + "provider_name": provider.name, + "provider_type": provider.provider_type, + }), + ) + .await; + Ok((StatusCode::CREATED, Json(response))) } @@ -600,7 +637,7 @@ async fn get_dns_provider( State(state): State>, Path(id): Path, ) -> Result { - permission_check!(auth, Permission::SettingsRead); + permission_check!(auth, Permission::DnsProvidersRead); let provider = state.provider_service.get(id).await?; @@ -650,11 +687,18 @@ async fn get_dns_provider( async fn update_provider( RequireAuth(auth): RequireAuth, State(state): State>, + Extension(metadata): Extension, Path(id): Path, Json(request): Json, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); - + permission_check!(auth, Permission::DnsProvidersWrite); + + let changed_fields = serde_json::json!({ + "name": request.name.is_some(), + "credentials": request.credentials.is_some(), + "description": request.description.is_some(), + "is_active": request.is_active, + }); let credentials: Option = request.credentials.map(|c| c.into()); if let Some(credentials) = &credentials { @@ -699,6 +743,17 @@ async fn update_provider( updated_at: provider.updated_at.to_rfc3339(), }; + log_dns_governance_audit( + &state, + &auth, + &metadata, + provider.id, + "", + "DNS_PROVIDER_UPDATED", + changed_fields, + ) + .await; + Ok(Json(response)) } @@ -718,11 +773,26 @@ async fn update_provider( async fn delete_dns_provider( RequireAuth(auth): RequireAuth, State(state): State>, + Extension(metadata): Extension, Path(id): Path, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); + let provider = state.provider_service.get(id).await?; state.provider_service.delete(id).await?; + log_dns_governance_audit( + &state, + &auth, + &metadata, + id, + "", + "DNS_PROVIDER_DELETED", + serde_json::json!({ + "provider_name": provider.name, + "provider_type": provider.provider_type, + }), + ) + .await; Ok(StatusCode::NO_CONTENT) } @@ -743,11 +813,22 @@ async fn delete_dns_provider( async fn test_provider_connection( RequireAuth(auth): RequireAuth, State(state): State>, + Extension(metadata): Extension, Path(id): Path, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); let success = state.provider_service.test_connection(id).await?; + log_dns_governance_audit( + &state, + &auth, + &metadata, + id, + "", + "DNS_PROVIDER_CONNECTION_TESTED", + serde_json::json!({ "success": success }), + ) + .await; let response = ConnectionTestResult { success, @@ -779,7 +860,7 @@ async fn list_provider_zones( State(state): State>, Path(id): Path, ) -> Result { - permission_check!(auth, Permission::SettingsRead); + permission_check!(auth, Permission::DnsProvidersRead); let provider = state.provider_service.get(id).await?; let instance = state.provider_service.create_provider_instance(&provider)?; @@ -807,10 +888,14 @@ async fn list_provider_zones( async fn add_managed_domain( RequireAuth(auth): RequireAuth, State(state): State>, + Extension(metadata): Extension, Path(id): Path, Json(request): Json, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); + if managed_domain_automation_enabled(request.auto_manage, request.sync_generated_records) { + permission_check!(auth, Permission::DnsAutomationWrite); + } let managed = state .provider_service @@ -825,6 +910,20 @@ async fn add_managed_domain( ) .await?; + log_dns_governance_audit( + &state, + &auth, + &metadata, + id, + &managed.domain, + "DNS_MANAGED_DOMAIN_ADDED", + serde_json::json!({ + "auto_manage": managed.auto_manage, + "verified": managed.verified, + }), + ) + .await; + Ok(( StatusCode::CREATED, Json(ManagedDomainResponse::from(managed)), @@ -849,7 +948,7 @@ async fn list_managed_domains( State(state): State>, Path(id): Path, ) -> Result { - permission_check!(auth, Permission::SettingsRead); + permission_check!(auth, Permission::DnsProvidersRead); let domains = state.provider_service.list_managed_domains(id).await?; @@ -877,14 +976,25 @@ async fn list_managed_domains( async fn remove_managed_domain( RequireAuth(auth): RequireAuth, State(state): State>, + Extension(metadata): Extension, Path((provider_id, domain)): Path<(i32, String)>, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); state .provider_service .remove_managed_domain(provider_id, &domain) .await?; + log_dns_governance_audit( + &state, + &auth, + &metadata, + provider_id, + &domain, + "DNS_MANAGED_DOMAIN_REMOVED", + serde_json::json!({}), + ) + .await; Ok(StatusCode::NO_CONTENT) } @@ -905,9 +1015,10 @@ async fn remove_managed_domain( async fn verify_managed_domain( RequireAuth(auth): RequireAuth, State(state): State>, + Extension(metadata): Extension, Path((provider_id, domain)): Path<(i32, String)>, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); let _verified = state .provider_service @@ -923,6 +1034,16 @@ async fn verify_managed_domain( .into_iter() .find(|d| d.domain == domain) .ok_or_else(|| DnsError::DomainNotFound(domain))?; + log_dns_governance_audit( + &state, + &auth, + &metadata, + provider_id, + &managed.domain, + "DNS_MANAGED_DOMAIN_VERIFIED", + serde_json::json!({ "verified": managed.verified }), + ) + .await; Ok(Json(ManagedDomainResponse::from(managed))) } @@ -948,7 +1069,23 @@ async fn update_managed_domain( Path((provider_id, domain)): Path<(i32, String)>, Json(request): Json, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); + let existing = state + .provider_service + .get_managed_domain(provider_id, &domain) + .await?; + let resulting_auto_manage = request.auto_manage.unwrap_or(existing.auto_manage); + let resulting_sync_generated_records = request + .sync_generated_records + .unwrap_or(existing.sync_generated_records); + if managed_domain_automation_enabled(resulting_auto_manage, resulting_sync_generated_records) { + permission_check!(auth, Permission::DnsAutomationWrite); + } + let changes = serde_json::json!({ + "generated_hostname_mode": request.generated_hostname_mode, + "sync_generated_records": request.sync_generated_records, + "auto_manage": request.auto_manage, + }); let updated = state .provider_service @@ -963,13 +1100,14 @@ async fn update_managed_domain( ) .await?; - log_managed_domain_audit( + log_dns_governance_audit( &state, &auth, &metadata, provider_id, &domain, "DNS_MANAGED_DOMAIN_UPDATED", + changes, ) .await; @@ -1009,7 +1147,7 @@ async fn preview_hostname_mode( Path((provider_id, domain)): Path<(i32, String)>, Query(query): Query, ) -> Result { - permission_check!(auth, Permission::SettingsRead); + permission_check!(auth, Permission::DnsProvidersRead); let target = PublicHostnameStrategy::from_db_str(&query.mode); let result = state @@ -1042,7 +1180,10 @@ async fn apply_hostname_mode( Path((provider_id, domain)): Path<(i32, String)>, Json(request): Json, ) -> Result { - permission_check!(auth, Permission::SettingsWrite); + permission_check!(auth, Permission::DnsProvidersWrite); + if request.sync_dns { + permission_check!(auth, Permission::DnsAutomationWrite); + } let target = PublicHostnameStrategy::from_db_str(&request.mode); let result = state @@ -1066,13 +1207,17 @@ async fn apply_hostname_mode( ); } - log_managed_domain_audit( + log_dns_governance_audit( &state, &auth, &metadata, provider_id, &domain, "DNS_HOSTNAME_MODE_APPLIED", + serde_json::json!({ + "mode": request.mode, + "sync_dns": request.sync_dns, + }), ) .await; @@ -1080,15 +1225,16 @@ async fn apply_hostname_mode( } /// Emit an audit log for a managed-domain write; failure is logged, not fatal. -async fn log_managed_domain_audit( +async fn log_dns_governance_audit( state: &Arc, auth: &temps_auth::AuthContext, metadata: &RequestMetadata, provider_id: i32, domain: &str, action: &str, + details: serde_json::Value, ) { - let audit = ManagedDomainAudit { + let audit = DnsGovernanceAudit { context: AuditContext { user_id: auth.user_id(), ip_address: Some(metadata.ip_address.clone()), @@ -1097,6 +1243,7 @@ async fn log_managed_domain_audit( provider_id, domain: domain.to_string(), action: action.to_string(), + details, }; if let Err(e) = state.audit.create_audit_log(&audit).await { tracing::error!("Failed to create audit log: {}", e); @@ -1218,3 +1365,18 @@ pub fn configure_internal_routes() -> Router> { ) )] pub struct DnsApiDoc; + +#[cfg(test)] +mod tests { + use super::managed_domain_automation_enabled; + + #[test] + fn generated_record_sync_is_an_automation_capability() { + assert!(managed_domain_automation_enabled(false, true)); + } + + #[test] + fn manual_domain_without_generated_sync_is_not_automation() { + assert!(!managed_domain_automation_enabled(false, false)); + } +} diff --git a/crates/temps-dns/src/providers/namecheap.rs b/crates/temps-dns/src/providers/namecheap.rs index 69b2c9dfd..1c612cf8c 100644 --- a/crates/temps-dns/src/providers/namecheap.rs +++ b/crates/temps-dns/src/providers/namecheap.rs @@ -29,12 +29,30 @@ pub struct NamecheapProvider { } impl NamecheapProvider { + /// Summarize an API request without exposing parameter values in logs. + fn api_request_summary(command: &str, params: &[(&str, &str)]) -> String { + let parameter_names = params + .iter() + .map(|(name, _)| *name) + .collect::>() + .join(","); + + format!( + "command={} parameter_count={} parameter_names=[{}]", + command, + params.len(), + parameter_names + ) + } + /// Create a new Namecheap provider with the given credentials pub fn new(credentials: NamecheapCredentials) -> Result { let client = Client::builder() .timeout(std::time::Duration::from_secs(30)) .build() - .map_err(|e| DnsError::ApiError(format!("Failed to create HTTP client: {}", e)))?; + .map_err(|e| { + DnsError::ApiError(format!("Failed to create HTTP client: {}", e.without_url())) + })?; let base_url = if credentials.sandbox { NAMECHEAP_SANDBOX_URL.to_string() @@ -61,12 +79,13 @@ impl NamecheapProvider { .get("https://api.ipify.org") .send() .await - .map_err(|e| DnsError::ApiError(format!("Failed to get public IP: {}", e)))?; + .map_err(|e| { + DnsError::ApiError(format!("Failed to get public IP: {}", e.without_url())) + })?; - response - .text() - .await - .map_err(|e| DnsError::ApiError(format!("Failed to read IP response: {}", e))) + response.text().await.map_err(|e| { + DnsError::ApiError(format!("Failed to read IP response: {}", e.without_url())) + }) } /// Make an API request to Namecheap @@ -88,8 +107,8 @@ impl NamecheapProvider { query_params.extend(params); debug!( - "Namecheap API request: {} with params: {:?}", - command, params + request = %Self::api_request_summary(command, params), + "Namecheap API request" ); let response = self @@ -98,13 +117,12 @@ impl NamecheapProvider { .query(&query_params) .send() .await - .map_err(|e| DnsError::ApiError(format!("API request failed: {}", e)))?; + .map_err(|e| DnsError::ApiError(format!("API request failed: {}", e.without_url())))?; let status = response.status(); - let body = response - .text() - .await - .map_err(|e| DnsError::ApiError(format!("Failed to read response: {}", e)))?; + let body = response.text().await.map_err(|e| { + DnsError::ApiError(format!("Failed to read response: {}", e.without_url())) + })?; if !status.is_success() { return Err(DnsError::ApiError(format!( @@ -518,6 +536,24 @@ mod tests { assert_eq!(tld, "com"); } + #[test] + fn test_api_request_summary_excludes_parameter_values() { + const SENTINEL_TXT_VALUE: &str = "acme-secret-sentinel-value"; + let params = [ + ("SLD", "example"), + ("TLD", "com"), + ("Address1", SENTINEL_TXT_VALUE), + ]; + + let summary = + NamecheapProvider::api_request_summary("namecheap.domains.dns.setHosts", ¶ms); + + assert!(!summary.contains(SENTINEL_TXT_VALUE)); + assert!(summary.contains("command=namecheap.domains.dns.setHosts")); + assert!(summary.contains("parameter_count=3")); + assert!(summary.contains("parameter_names=[SLD,TLD,Address1]")); + } + #[test] fn test_split_domain_co_uk() { // Note: Namecheap treats co.uk as sld.tld @@ -933,6 +969,34 @@ mod integration_tests { } } + #[tokio::test] + async fn test_list_zones_transport_failure_redacts_api_key() { + const SECRET_API_KEY: &str = "namecheap-secret-api-key-sentinel"; + let provider = NamecheapProvider { + client: Client::builder() + .timeout(std::time::Duration::from_secs(1)) + .build() + .unwrap(), + credentials: NamecheapCredentials { + api_user: "testuser".to_string(), + api_key: SECRET_API_KEY.to_string(), + client_ip: Some("127.0.0.1".to_string()), + sandbox: true, + }, + // Port 9 is expected to refuse the connection, forcing reqwest's + // transport-error path after query parameters have been attached. + base_url: "http://127.0.0.1:9/xml.response".to_string(), + }; + + let error = provider + .list_zones() + .await + .expect_err("the forced transport failure must return an error"); + + assert!(error.to_string().contains("API request failed")); + assert!(!error.to_string().contains(SECRET_API_KEY)); + } + #[tokio::test] async fn test_list_zones_success() { let mock_server = MockServer::start().await; diff --git a/crates/temps-dns/src/services/provider_service.rs b/crates/temps-dns/src/services/provider_service.rs index 9a7cb072b..e02345a04 100644 --- a/crates/temps-dns/src/services/provider_service.rs +++ b/crates/temps-dns/src/services/provider_service.rs @@ -8,7 +8,7 @@ use sea_orm::{ ActiveModelTrait, ActiveValue::Set, ColumnTrait, DatabaseConnection, EntityTrait, QueryFilter, - QueryOrder, + QueryOrder, QuerySelect, }; use std::sync::Arc; use temps_core::EncryptionService; @@ -69,6 +69,12 @@ pub struct UpdateManagedDomainRequest { } impl DnsProviderService { + // A valid DNS name has at most 127 labels. Keeping the candidate set capped + // also bounds malformed input before it reaches an IN predicate. + const MAX_AUTHORITATIVE_SUFFIX_CANDIDATES: usize = 127; + const NORMALIZED_MANAGED_DOMAIN_SQL: &'static str = + "LOWER(REGEXP_REPLACE(RTRIM(BTRIM(\"dns_managed_domains\".\"domain\"), '.'), '^((\\*\\.)+)', ''))"; + pub fn new(db: Arc, encryption_service: Arc) -> Self { Self { db, @@ -360,6 +366,13 @@ impl DnsProviderService { &self, provider: &dns_providers::Model, ) -> Result, DnsError> { + if !provider.is_active { + return Err(DnsError::ProviderInactive { + provider_id: provider.id, + provider_name: provider.name.clone(), + }); + } + // Decrypt credentials let credentials_json = self .encryption_service @@ -504,6 +517,13 @@ impl DnsProviderService { &self, provider: &dns_providers::Model, ) -> Result { + if !provider.is_active { + return Err(DnsError::ProviderInactive { + provider_id: provider.id, + provider_name: provider.name.clone(), + }); + } + let credentials_json = self .encryption_service .decrypt_string(&provider.credentials) @@ -527,17 +547,33 @@ impl DnsProviderService { // Verify provider exists let _provider = self.get(provider_id).await?; - // Check if domain is already managed + let canonical_domain = Self::normalize_domain(&request.domain); + if canonical_domain.is_empty() { + return Err(DnsError::Validation(format!( + "Managed domain '{}' has no DNS labels after canonicalization", + request.domain + ))); + } + + // Check the same canonical form used by authoritative lookup. New rows + // are stored canonically, so the existing raw unique constraint also + // closes the race between this preflight query and the insert. let existing = dns_managed_domains::Entity::find() - .filter(dns_managed_domains::Column::Domain.eq(&request.domain)) + .filter(sea_orm::sea_query::Expr::cust_with_values( + format!("{} = $1", Self::NORMALIZED_MANAGED_DOMAIN_SQL), + [canonical_domain.clone()], + )) + .limit(1) .one(self.db.as_ref()) .await?; - if existing.is_some() { - return Err(DnsError::Validation(format!( - "Domain {} is already managed by another provider", - request.domain - ))); + if let Some(existing) = existing { + return Err(DnsError::ManagedDomainAlreadyExists { + requested_domain: request.domain, + canonical_domain, + existing_managed_domain_id: existing.id, + existing_provider_id: existing.provider_id, + }); } // Normalize the requested mode; unknown values fall back to standard. @@ -552,7 +588,7 @@ impl DnsProviderService { let managed_domain = dns_managed_domains::ActiveModel { provider_id: Set(provider_id), - domain: Set(request.domain.clone()), + domain: Set(canonical_domain.clone()), auto_manage: Set(request.auto_manage), verified: Set(false), generated_hostname_mode: Set(mode), @@ -564,7 +600,7 @@ impl DnsProviderService { info!( "Added managed domain {} to provider {}", - request.domain, provider_id + canonical_domain, provider_id ); Ok(result) @@ -608,6 +644,21 @@ impl DnsProviderService { Ok(domains) } + /// Load one managed domain so authorization can be evaluated against the + /// resulting settings before a handler mutates it. + pub async fn get_managed_domain( + &self, + provider_id: i32, + domain: &str, + ) -> Result { + dns_managed_domains::Entity::find() + .filter(dns_managed_domains::Column::ProviderId.eq(provider_id)) + .filter(dns_managed_domains::Column::Domain.eq(domain)) + .one(self.db.as_ref()) + .await? + .ok_or_else(|| DnsError::DomainNotFound(domain.to_string())) + } + /// Verify a managed domain (check if provider can access it) pub async fn verify_managed_domain( &self, @@ -618,7 +669,7 @@ impl DnsProviderService { let instance = self.create_provider_instance(&provider)?; // Distinguish "token lacks zone access" (PermissionDenied) from "zone - // absent" so the UI can flag a mis-scoped token. + // absent" so the UI can flag an incorrectly scoped token. let access = instance.check_zone_access(domain).await; let can_manage = access.is_ok(); let (zone_access_ok, zone_access_error) = match &access { @@ -867,34 +918,135 @@ impl DnsProviderService { &self, domain: &str, ) -> Result, DnsError> { - // Extract base domain - let base_domain = Self::extract_base_domain(domain); + self.find_active_managed_domain_candidate(domain, None, true) + .await + } - let managed_domain = dns_managed_domains::Entity::find() - .filter(dns_managed_domains::Column::Domain.eq(&base_domain)) + /// Find the verified zone belonging to `provider_id` that authoritatively + /// covers `domain`. Human-triggered writes do not require `auto_manage`, but + /// they must never be allowed to use arbitrary provider credentials. + pub async fn find_verified_zone_for_provider( + &self, + provider_id: i32, + domain: &str, + ) -> Result, DnsError> { + Ok(self + .find_active_managed_domain_candidate(domain, Some(provider_id), false) + .await? + .map(|(_, managed)| managed)) + } + + /// Find the longest authoritative suffix without loading the managed-zone + /// table. The normalized equality predicate is backed by + /// `idx_dns_managed_domains_normalized_domain`, preserving legacy mixed-case, + /// wildcard, whitespace, and trailing-dot rows without a sequential scan. + async fn find_active_managed_domain_candidate( + &self, + domain: &str, + provider_id: Option, + require_auto_manage: bool, + ) -> Result, DnsError> { + let candidates = Self::authoritative_suffix_candidates(domain); + if candidates.is_empty() { + return Ok(None); + } + + let placeholders = (1..=candidates.len()) + .map(|position| format!("${position}")) + .collect::>() + .join(", "); + let mut query = dns_managed_domains::Entity::find() + .find_also_related(dns_providers::Entity) + .filter(sea_orm::sea_query::Expr::cust_with_values( + format!( + "{} IN ({placeholders})", + Self::NORMALIZED_MANAGED_DOMAIN_SQL + ), + candidates, + )) .filter(dns_managed_domains::Column::Verified.eq(true)) - .filter(dns_managed_domains::Column::AutoManage.eq(true)) - .one(self.db.as_ref()) - .await?; + // This predicate must execute in the same SQL statement before + // LIMIT, otherwise an inactive duplicate can shadow an active zone. + .filter(dns_providers::Column::IsActive.eq(true)); + if let Some(provider_id) = provider_id { + query = query.filter(dns_managed_domains::Column::ProviderId.eq(provider_id)); + } + if require_auto_manage { + query = query.filter(dns_managed_domains::Column::AutoManage.eq(true)); + } - if let Some(managed) = managed_domain { - let provider = self.get(managed.provider_id).await?; - if provider.is_active { - return Ok(Some((provider, managed))); - } + let rows = query + .order_by_desc(sea_orm::sea_query::Expr::cust(format!( + "CHAR_LENGTH({})", + Self::NORMALIZED_MANAGED_DOMAIN_SQL + ))) + .order_by_asc(dns_managed_domains::Column::Id) + .limit(2) + .all(self.db.as_ref()) + .await + .map_err(DnsError::from)? + .into_iter() + .filter_map(|(managed, provider)| provider.map(|provider| (provider, managed))) + .collect::>(); + + let Some((provider, managed)) = rows.first() else { + return Ok(None); + }; + let canonical_zone = Self::normalize_domain(&managed.domain); + + if rows.get(1).is_some_and(|(_, candidate)| { + Self::normalize_domain(&candidate.domain) == canonical_zone + }) { + return Err(DnsError::AmbiguousManagedDomain { + requested_domain: domain.to_string(), + canonical_zone, + managed_domain_ids: rows.iter().map(|(_, managed)| managed.id).collect(), + provider_ids: rows.iter().map(|(provider, _)| provider.id).collect(), + }); } - Ok(None) + let mut selected = managed.clone(); + selected.domain = canonical_zone; + Ok(Some((provider.clone(), selected))) } - /// Extract base domain from a full domain name - fn extract_base_domain(domain: &str) -> String { - let parts: Vec<&str> = domain.split('.').collect(); - if parts.len() >= 2 { - parts[parts.len() - 2..].join(".") - } else { - domain.to_string() + fn authoritative_suffix_candidates(domain: &str) -> Vec { + let normalized = Self::normalize_domain(domain); + let mut candidates = Vec::new(); + let mut suffix = normalized.as_str(); + + while !suffix.is_empty() && candidates.len() < Self::MAX_AUTHORITATIVE_SUFFIX_CANDIDATES { + candidates.push(suffix.to_string()); + suffix = match suffix.find('.') { + Some(separator) => &suffix[separator + 1..], + None => break, + }; } + + candidates + } + + #[cfg(test)] + fn longest_managed_domain_match( + domain: &str, + managed_domains: Vec, + ) -> Option { + let domain = Self::normalize_domain(domain); + managed_domains + .into_iter() + .filter(|managed| { + let zone = Self::normalize_domain(&managed.domain); + domain == zone || domain.ends_with(&format!(".{zone}")) + }) + .max_by_key(|managed| Self::normalize_domain(&managed.domain).len()) + } + + fn normalize_domain(domain: &str) -> String { + domain + .trim() + .trim_start_matches("*.") + .trim_end_matches('.') + .to_ascii_lowercase() } } @@ -914,20 +1066,305 @@ impl temps_core::PublicHostnameResolver for DnsProviderService { #[cfg(test)] mod tests { use super::*; + use sea_orm::{DatabaseBackend, MockDatabase}; + + fn managed_domain(id: i32, provider_id: i32, domain: &str) -> dns_managed_domains::Model { + let now = chrono::Utc::now(); + dns_managed_domains::Model { + id, + provider_id, + domain: domain.to_string(), + zone_id: None, + auto_manage: true, + verified: true, + verified_at: Some(now), + verification_error: None, + generated_hostname_mode: "standard".to_string(), + sync_generated_records: false, + zone_access_ok: Some(true), + zone_access_error: None, + created_at: now, + updated_at: now, + } + } + + fn dns_provider(id: i32, name: &str, is_active: bool) -> dns_providers::Model { + let now = chrono::Utc::now(); + dns_providers::Model { + id, + name: name.to_string(), + provider_type: "cloudflare".to_string(), + credentials: "deliberately-not-encrypted".to_string(), + is_active, + description: None, + last_used_at: None, + last_error: None, + created_at: now, + updated_at: now, + } + } + + #[test] + fn inactive_provider_is_rejected_before_credentials_are_decrypted() { + let db = Arc::new(MockDatabase::new(DatabaseBackend::Postgres).into_connection()); + let service = DnsProviderService::new( + db, + Arc::new(EncryptionService::new_from_password( + "inactive-provider-test", + )), + ); + + let result = + service.create_provider_instance(&dns_provider(42, "disabled-cloudflare", false)); + + assert!(matches!( + result, + Err(DnsError::ProviderInactive { + provider_id: 42, + provider_name + }) if provider_name == "disabled-cloudflare" + )); + } #[test] - fn test_extract_base_domain() { + fn longest_managed_domain_suffix_wins() { + let matched = DnsProviderService::longest_managed_domain_match( + "API.Dev.Example.COM.", + vec![ + managed_domain(1, 10, "example.com"), + managed_domain(2, 20, "dev.example.com"), + ], + ); + assert_eq!( - DnsProviderService::extract_base_domain("example.com"), - "example.com" + matched.map(|managed| (managed.provider_id, managed.domain)), + Some((20, "dev.example.com".to_string())) ); + } + + #[test] + fn managed_domain_match_respects_label_boundaries_and_wildcards() { + let domains = vec![managed_domain(1, 10, "example.com")]; assert_eq!( - DnsProviderService::extract_base_domain("sub.example.com"), - "example.com" + DnsProviderService::longest_managed_domain_match("*.www.example.com", domains.clone()) + .map(|managed| managed.domain), + Some("example.com".to_string()) ); + assert!( + DnsProviderService::longest_managed_domain_match("notexample.com", domains).is_none() + ); + } + + #[test] + fn authoritative_suffix_candidates_are_normalized_longest_first() { + assert_eq!( + DnsProviderService::authoritative_suffix_candidates(" *.API.Dev.Example.COM. "), + vec![ + "api.dev.example.com", + "dev.example.com", + "example.com", + "com" + ] + ); + } + + #[test] + fn managed_domain_canonicalization_strips_repeated_wildcards_and_root_dots() { assert_eq!( - DnsProviderService::extract_base_domain("deep.sub.example.com"), - "example.com" + DnsProviderService::normalize_domain(" *.*.*.API.Example.COM... "), + "api.example.com" + ); + } + + #[test] + fn authoritative_suffix_candidates_are_bounded_for_malformed_input() { + let domain = std::iter::repeat_n("label", 200) + .collect::>() + .join("."); + + let candidates = DnsProviderService::authoritative_suffix_candidates(&domain); + + assert_eq!( + candidates.len(), + DnsProviderService::MAX_AUTHORITATIVE_SUFFIX_CANDIDATES + ); + assert_eq!(candidates.first(), Some(&domain)); + } + + #[tokio::test] + async fn verified_zone_query_binds_candidates_before_provider_filter() { + let db = Arc::new( + MockDatabase::new(DatabaseBackend::Postgres) + .append_query_results(vec![Vec::::new()]) + .into_connection(), + ); + let service = DnsProviderService::new( + db.clone(), + Arc::new(EncryptionService::new_from_password("provider-query-test")), + ); + + let result = service + .find_verified_zone_for_provider(42, "*.API.Dev.Example.COM.") + .await; + + assert!(matches!(result, Ok(None))); + drop(service); + let transaction_log = Arc::try_unwrap(db) + .expect("test must release the database connection") + .into_transaction_log(); + let statement = &transaction_log[0].statements()[0]; + assert!(statement.sql.contains( + "LOWER(REGEXP_REPLACE(RTRIM(BTRIM(\"dns_managed_domains\".\"domain\"), '.'), '^((\\*\\.)+)', '')) IN ($1, $2, $3, $4)" + )); + assert!(statement.sql.contains("LEFT JOIN \"dns_providers\"")); + assert!(statement + .sql + .contains(r#""dns_providers"."is_active" = $6"#)); + assert!(statement + .sql + .contains(r#""dns_managed_domains"."provider_id" = $7"#)); + assert!(statement.sql.contains("LIMIT $8")); + assert_eq!( + format!("{:?}", statement.values), + "Some(Values([String(Some(\"api.dev.example.com\")), String(Some(\"dev.example.com\")), String(Some(\"example.com\")), String(Some(\"com\")), Bool(Some(true)), Bool(Some(true)), Int(Some(42)), BigUnsigned(Some(2))]))" + ); + } + + #[tokio::test] + async fn duplicate_best_canonical_zone_fails_closed() { + let db = Arc::new( + MockDatabase::new(DatabaseBackend::Postgres) + .append_query_results(vec![vec![ + ( + managed_domain(11, 101, "Example.COM."), + Some(dns_provider(101, "first", true)), + ), + ( + managed_domain(12, 202, "*.example.com"), + Some(dns_provider(202, "second", true)), + ), + ]]) + .into_connection(), + ); + let service = DnsProviderService::new( + db, + Arc::new(EncryptionService::new_from_password("ambiguous-zone-test")), + ); + + let result = service.find_provider_for_domain("api.example.com").await; + + assert!(matches!( + result, + Err(DnsError::AmbiguousManagedDomain { + requested_domain, + canonical_zone, + managed_domain_ids, + provider_ids, + }) if requested_domain == "api.example.com" + && canonical_zone == "example.com" + && managed_domain_ids == vec![11, 12] + && provider_ids == vec![101, 202] + )); + } + + #[tokio::test] + async fn more_specific_zone_wins_without_treating_parent_as_ambiguous() { + let child = managed_domain(11, 101, "*.Dev.Example.COM."); + let db = Arc::new( + MockDatabase::new(DatabaseBackend::Postgres) + .append_query_results(vec![vec![ + (child.clone(), Some(dns_provider(101, "child", true))), + ( + managed_domain(12, 202, "example.com"), + Some(dns_provider(202, "parent", true)), + ), + ]]) + .into_connection(), + ); + let service = DnsProviderService::new( + db, + Arc::new(EncryptionService::new_from_password("parent-zone-test")), + ); + + let result = service + .find_provider_for_domain("api.dev.example.com") + .await; + + assert!(matches!( + result, + Ok(Some((provider, managed))) if provider.id == 101 + && managed.id == child.id + && managed.domain == "dev.example.com" + )); + } + + #[tokio::test] + async fn add_managed_domain_rejects_canonical_duplicate_using_indexed_predicate() { + let db = Arc::new( + MockDatabase::new(DatabaseBackend::Postgres) + .append_query_results(vec![vec![dns_provider(101, "provider", true)]]) + .append_query_results(vec![vec![managed_domain(11, 202, "example.com")]]) + .into_connection(), + ); + let service = DnsProviderService::new( + db.clone(), + Arc::new(EncryptionService::new_from_password("canonical-add-test")), + ); + + let result = service + .add_managed_domain( + 101, + AddManagedDomainRequest { + domain: " *.*.Example.COM... ".to_string(), + auto_manage: true, + generated_hostname_mode: None, + sync_generated_records: false, + }, + ) + .await; + + assert!(matches!( + result, + Err(DnsError::ManagedDomainAlreadyExists { + requested_domain, + canonical_domain, + existing_managed_domain_id: 11, + existing_provider_id: 202, + }) if requested_domain == " *.*.Example.COM... " + && canonical_domain == "example.com" + )); + drop(service); + let transaction_log = Arc::try_unwrap(db) + .expect("test must release the database connection") + .into_transaction_log(); + let statement = &transaction_log[1].statements()[0]; + assert!(statement.sql.contains( + "LOWER(REGEXP_REPLACE(RTRIM(BTRIM(\"dns_managed_domains\".\"domain\"), '.'), '^((\\*\\.)+)', '')) = $1" + )); + assert!(statement.sql.contains("LIMIT $2")); + assert_eq!( + format!("{:?}", statement.values), + "Some(Values([String(Some(\"example.com\")), BigUnsigned(Some(1))]))" + ); + } + + #[test] + fn masked_credentials_for_inactive_provider_fail_before_decryption() { + let service = DnsProviderService::new( + Arc::new(MockDatabase::new(DatabaseBackend::Postgres).into_connection()), + Arc::new(EncryptionService::new_from_password("inactive-mask-test")), ); + let mut provider = dns_provider(101, "inactive", false); + provider.credentials = "not-valid-ciphertext".to_string(); + + let result = service.get_masked_credentials(&provider); + + assert!(matches!( + result, + Err(DnsError::ProviderInactive { + provider_id: 101, + provider_name, + }) if provider_name == "inactive" + )); } } diff --git a/crates/temps-dns/tests/governance_router_test.rs b/crates/temps-dns/tests/governance_router_test.rs new file mode 100644 index 000000000..ae6c4fca5 --- /dev/null +++ b/crates/temps-dns/tests/governance_router_test.rs @@ -0,0 +1,711 @@ +use std::sync::{Arc, Mutex}; + +use async_trait::async_trait; +use axum::{ + body::Body, + http::{Method, Request, StatusCode}, +}; +use sea_orm::{DatabaseBackend, MockDatabase}; +use temps_auth::{AuthContext, Permission}; +use temps_core::{AuditLogger, AuditOperation, Job, JobQueue, JobReceiver, RequestMetadata}; +use temps_dns::{ + handlers::{configure_routes, DnsAppState}, + services::{DnsProviderService, DnsRecordService}, +}; +use tower::ServiceExt; + +struct NoopQueue; + +#[async_trait] +impl JobQueue for NoopQueue { + async fn send(&self, _job: Job) -> Result<(), temps_core::QueueError> { + Ok(()) + } + + fn subscribe(&self) -> Box { + panic!("permission-boundary tests never subscribe") + } +} + +struct NoopAudit; + +#[async_trait] +impl AuditLogger for NoopAudit { + async fn create_audit_log(&self, _operation: &dyn AuditOperation) -> anyhow::Result<()> { + Ok(()) + } +} + +#[derive(Default)] +struct RecordingAudit { + operations: Mutex>, +} + +#[async_trait] +impl AuditLogger for RecordingAudit { + async fn create_audit_log(&self, operation: &dyn AuditOperation) -> anyhow::Result<()> { + self.operations + .lock() + .unwrap() + .push(operation.operation_type()); + Ok(()) + } +} + +fn test_user() -> temps_entities::users::Model { + let now = chrono::Utc::now(); + temps_entities::users::Model { + id: 42, + name: "DNS operator".to_string(), + email: "dns@example.com".to_string(), + password_hash: None, + email_verified: true, + email_verification_token: None, + email_verification_expires: None, + password_reset_token: None, + password_reset_expires: None, + must_change_password: false, + deleted_at: None, + mfa_secret: None, + mfa_enabled: false, + mfa_recovery_codes: None, + oidc_subject: None, + oidc_provider_id: None, + created_at: now, + updated_at: now, + } +} + +fn auth(permissions: Vec) -> AuthContext { + AuthContext::new_api_key( + test_user(), + None, + Some(permissions), + "governance-test".to_string(), + 1, + ) +} + +fn metadata() -> RequestMetadata { + RequestMetadata { + ip_address: "127.0.0.1".to_string(), + user_agent: "governance-router-test".to_string(), + headers: Default::default(), + visitor_id_cookie: None, + session_id_cookie: None, + base_url: "http://localhost".to_string(), + scheme: "http".to_string(), + host: "localhost".to_string(), + is_secure: false, + } +} + +fn router() -> axum::Router { + // No query results are registered. A permission regression that touches the + // service/DB before returning 403 therefore fails loudly instead of merely + // returning the same status for the wrong reason. + let db = Arc::new(MockDatabase::new(DatabaseBackend::Postgres).into_connection()); + let provider_service = Arc::new(DnsProviderService::new( + db, + Arc::new(temps_core::EncryptionService::new_from_password("test")), + )); + let state = Arc::new(DnsAppState { + record_service: Arc::new(DnsRecordService::new(provider_service.clone())), + provider_service, + queue: Arc::new(NoopQueue), + audit: Arc::new(NoopAudit), + }); + configure_routes().with_state(state) +} + +fn router_with_db( + db: Arc, + encryption: Arc, +) -> axum::Router { + let provider_service = Arc::new(DnsProviderService::new(db, encryption)); + let state = Arc::new(DnsAppState { + record_service: Arc::new(DnsRecordService::new(provider_service.clone())), + provider_service, + queue: Arc::new(NoopQueue), + audit: Arc::new(NoopAudit), + }); + configure_routes().with_state(state) +} + +fn request_for( + method: Method, + uri: impl AsRef, + permissions: Vec, + body: impl Into, +) -> Request { + let mut request = Request::builder() + .method(method) + .uri(uri.as_ref()) + .header("content-type", "application/json") + .body(body.into()) + .unwrap(); + request.extensions_mut().insert(auth(permissions)); + request.extensions_mut().insert(metadata()); + request +} + +async fn request( + method: Method, + uri: &str, + permissions: Vec, + body: &str, +) -> StatusCode { + let mut request = Request::builder() + .method(method) + .uri(uri) + .header("content-type", "application/json") + .body(Body::from(body.to_string())) + .unwrap(); + request.extensions_mut().insert(auth(permissions)); + request.extensions_mut().insert(metadata()); + router().oneshot(request).await.unwrap().status() +} + +#[tokio::test] +async fn test_list_dns_providers_without_read_permission_returns_forbidden_before_db_touch() { + let status = request( + Method::GET, + "/dns-providers", + vec![Permission::DnsProvidersWrite], + "", + ) + .await; + + assert_eq!(status, StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn test_remove_managed_domain_without_write_permission_returns_forbidden_before_db_touch() { + let status = request( + Method::DELETE, + "/dns-providers/7/domains/example.com", + vec![Permission::DnsProvidersRead], + "", + ) + .await; + + assert_eq!(status, StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn test_add_auto_managed_domain_without_automation_permission_returns_forbidden_before_db_touch( +) { + let status = request( + Method::POST, + "/dns-providers/7/domains", + vec![Permission::DnsProvidersWrite], + r#"{"domain":"example.com","auto_manage":true,"sync_generated_records":false}"#, + ) + .await; + + assert_eq!(status, StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn test_add_sync_enabled_domain_without_automation_permission_returns_forbidden_before_db_touch( +) { + let status = request( + Method::POST, + "/dns-providers/7/domains", + vec![Permission::DnsProvidersWrite], + r#"{"domain":"example.com","auto_manage":false,"sync_generated_records":true}"#, + ) + .await; + + assert_eq!(status, StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn test_add_non_automated_domain_does_not_require_automation_permission() { + let status = request( + Method::POST, + "/dns-providers/7/domains", + vec![Permission::DnsProvidersWrite], + r#"{"domain":"example.com","auto_manage":false,"sync_generated_records":false}"#, + ) + .await; + + assert_ne!( + status, + StatusCode::FORBIDDEN, + "dns:providers:write alone must authorize auto_manage=false; the empty mock DB may fail later" + ); +} + +#[tokio::test] +async fn test_apply_hostname_mode_with_dns_sync_without_automation_permission_returns_forbidden_before_db_touch( +) { + let status = request( + Method::POST, + "/dns-providers/7/domains/example.com/apply-hostname-mode", + vec![Permission::DnsProvidersWrite], + r#"{"mode":"flat","sync_dns":true}"#, + ) + .await; + + assert_eq!(status, StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn test_apply_hostname_mode_without_dns_sync_does_not_require_automation_permission() { + let status = request( + Method::POST, + "/dns-providers/7/domains/example.com/apply-hostname-mode", + vec![Permission::DnsProvidersWrite], + r#"{"mode":"flat","sync_dns":false}"#, + ) + .await; + + assert_ne!( + status, + StatusCode::FORBIDDEN, + "dns:providers:write must retain access when sync_dns is false; the empty mock DB may fail later" + ); +} + +#[tokio::test] +async fn test_add_managed_domain_success_emits_governance_audit() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_entities::dns_providers; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping DNS router audit test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let provider = dns_providers::ActiveModel { + name: Set("manual-dns".to_string()), + provider_type: Set("manual".to_string()), + credentials: Set(encryption.encrypt_string("{}").unwrap()), + is_active: Set(true), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert provider"); + let provider_service = Arc::new(DnsProviderService::new(db, encryption)); + let audit = Arc::new(RecordingAudit::default()); + let state = Arc::new(DnsAppState { + record_service: Arc::new(DnsRecordService::new(provider_service.clone())), + provider_service, + queue: Arc::new(NoopQueue), + audit: audit.clone(), + }); + let mut request = Request::builder() + .method(Method::POST) + .uri(format!("/dns-providers/{}/domains", provider.id)) + .header("content-type", "application/json") + .body(Body::from( + r#"{"domain":"example.com","auto_manage":false,"sync_generated_records":false}"#, + )) + .unwrap(); + request + .extensions_mut() + .insert(auth(vec![Permission::DnsProvidersWrite])); + request.extensions_mut().insert(metadata()); + + let response = configure_routes() + .with_state(state) + .oneshot(request) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::CREATED); + assert_eq!( + audit.operations.lock().unwrap().as_slice(), + ["DNS_MANAGED_DOMAIN_ADDED"] + ); +} + +#[tokio::test] +async fn test_update_existing_automated_domain_enabling_sync_requires_automation_permission() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_entities::{dns_managed_domains, dns_providers}; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping DNS update permission test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let provider = dns_providers::ActiveModel { + name: Set("manual-dns".to_string()), + provider_type: Set("manual".to_string()), + credentials: Set(encryption.encrypt_string("{}").unwrap()), + is_active: Set(true), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert provider"); + dns_managed_domains::ActiveModel { + provider_id: Set(provider.id), + domain: Set("example.com".to_string()), + auto_manage: Set(true), + verified: Set(true), + generated_hostname_mode: Set("standard".to_string()), + sync_generated_records: Set(false), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert managed domain"); + let uri = format!("/dns-providers/{}/domains/example.com", provider.id); + + let forbidden = router_with_db(db.clone(), encryption.clone()) + .oneshot(request_for( + Method::PATCH, + &uri, + vec![Permission::DnsProvidersWrite], + Body::from(r#"{"sync_generated_records":true}"#), + )) + .await + .unwrap(); + assert_eq!(forbidden.status(), StatusCode::FORBIDDEN); + + let authorized = router_with_db(db, encryption) + .oneshot(request_for( + Method::PATCH, + &uri, + vec![ + Permission::DnsProvidersWrite, + Permission::DnsAutomationWrite, + ], + Body::from(r#"{"sync_generated_records":true}"#), + )) + .await + .unwrap(); + assert_eq!(authorized.status(), StatusCode::OK); + let body = axum::body::to_bytes(authorized.into_body(), usize::MAX) + .await + .unwrap(); + let response: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!(response["auto_manage"], true); + assert_eq!(response["sync_generated_records"], true); +} + +#[tokio::test] +async fn test_list_zones_for_inactive_provider_rejects_before_credentials_are_decrypted() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_entities::dns_providers; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping inactive-provider zones test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let provider = dns_providers::ActiveModel { + name: Set("disabled-cloudflare".to_string()), + provider_type: Set("cloudflare".to_string()), + credentials: Set("not-valid-ciphertext".to_string()), + is_active: Set(false), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert inactive provider"); + + let response = router_with_db(db, encryption) + .oneshot(request_for( + Method::GET, + format!("/dns-providers/{}/zones", provider.id), + vec![Permission::DnsProvidersRead], + Body::empty(), + )) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + let problem: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!(problem["title"], "DNS Provider Is Inactive"); + assert!(problem["detail"] + .as_str() + .unwrap() + .contains("disabled-cloudflare")); +} + +#[tokio::test] +async fn test_find_provider_for_duplicate_zone_skips_inactive_provider_candidate() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_entities::{dns_managed_domains, dns_providers}; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping duplicate provider candidate test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let mut provider_ids = Vec::new(); + for (name, domain, is_active) in [ + ("inactive-provider", " EXAMPLE.COM. ", false), + ("active-provider", "example.com", true), + ] { + let provider = dns_providers::ActiveModel { + name: Set(name.to_string()), + provider_type: Set("manual".to_string()), + credentials: Set(encryption.encrypt_string("{}").unwrap()), + is_active: Set(is_active), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert provider"); + dns_managed_domains::ActiveModel { + provider_id: Set(provider.id), + domain: Set(domain.to_string()), + auto_manage: Set(true), + verified: Set(true), + generated_hostname_mode: Set("standard".to_string()), + sync_generated_records: Set(false), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert duplicate managed domain"); + provider_ids.push((provider.id, is_active)); + } + let active_provider_id = provider_ids + .iter() + .find_map(|(id, active)| active.then_some(*id)) + .unwrap(); + let service = DnsProviderService::new(db, encryption); + + let (provider, managed) = service + .find_provider_for_domain("app.example.com") + .await + .expect("provider lookup") + .expect("active provider candidate"); + + assert_eq!(provider.id, active_provider_id); + assert!(provider.is_active); + assert_eq!(managed.provider_id, active_provider_id); +} + +#[tokio::test] +#[serial_test::serial(dns_governance_db)] +async fn test_add_managed_domain_canonical_duplicate_returns_conflict() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_dns::services::AddManagedDomainRequest; + use temps_entities::{dns_managed_domains, dns_providers}; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping canonical duplicate router test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let encrypted_credentials = encryption.encrypt_string("{}").unwrap(); + + let existing_provider = dns_providers::ActiveModel { + name: Set("existing-provider".to_string()), + provider_type: Set("manual".to_string()), + credentials: Set(encrypted_credentials.clone()), + is_active: Set(true), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert existing provider"); + dns_managed_domains::ActiveModel { + provider_id: Set(existing_provider.id), + domain: Set(" *.EXAMPLE.COM. ".to_string()), + auto_manage: Set(false), + verified: Set(true), + generated_hostname_mode: Set("standard".to_string()), + sync_generated_records: Set(false), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert legacy non-canonical managed domain"); + let target_provider = dns_providers::ActiveModel { + name: Set("target-provider".to_string()), + provider_type: Set("manual".to_string()), + credentials: Set(encrypted_credentials), + is_active: Set(true), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert target provider"); + + let response = router_with_db(db.clone(), encryption.clone()) + .oneshot(request_for( + Method::POST, + format!("/dns-providers/{}/domains", target_provider.id), + vec![Permission::DnsProvidersWrite], + Body::from( + r#"{"domain":"example.com","auto_manage":false,"sync_generated_records":false}"#, + ), + )) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::CONFLICT); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + let problem: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!(problem["title"], "Managed DNS Domain Already Exists"); + assert!(problem["detail"] + .as_str() + .unwrap() + .contains("canonicalizes to 'example.com', which is already managed")); + + let service = DnsProviderService::new(db, encryption); + let fresh = service + .add_managed_domain( + target_provider.id, + AddManagedDomainRequest { + domain: " *.Fresh.Example.NET. ".to_string(), + auto_manage: false, + generated_hostname_mode: None, + sync_generated_records: false, + }, + ) + .await + .expect("add a fresh non-canonical managed domain"); + assert_eq!(fresh.domain, "fresh.example.net"); +} + +#[tokio::test] +#[serial_test::serial(dns_governance_db)] +async fn test_find_provider_for_canonical_duplicate_eligible_zones_returns_ambiguity() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_dns::errors::DnsError; + use temps_entities::{dns_managed_domains, dns_providers}; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping ambiguous managed-zone test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let encrypted_credentials = encryption.encrypt_string("{}").unwrap(); + let mut longest_zones = Vec::new(); + + for (name, domain) in [ + ("longest-one", "api.example.com"), + ("longest-two", " *.API.EXAMPLE.COM. "), + ("parent", "example.com"), + ] { + let provider = dns_providers::ActiveModel { + name: Set(name.to_string()), + provider_type: Set("manual".to_string()), + credentials: Set(encrypted_credentials.clone()), + is_active: Set(true), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert active provider"); + let managed = dns_managed_domains::ActiveModel { + provider_id: Set(provider.id), + domain: Set(domain.to_string()), + auto_manage: Set(true), + verified: Set(true), + generated_hostname_mode: Set("standard".to_string()), + sync_generated_records: Set(false), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert eligible managed domain"); + if name.starts_with("longest") { + longest_zones.push((provider.id, managed)); + } + } + let service = DnsProviderService::new(db.clone(), encryption); + + let error = service + .find_provider_for_domain("app.api.example.com") + .await + .expect_err("equivalent longest eligible zones must fail closed"); + match error { + DnsError::AmbiguousManagedDomain { + requested_domain, + canonical_zone, + managed_domain_ids, + provider_ids, + } => { + assert_eq!(requested_domain, "app.api.example.com"); + assert_eq!(canonical_zone, "api.example.com"); + assert_eq!(managed_domain_ids.len(), 2); + assert_eq!(provider_ids.len(), 2); + assert!(longest_zones + .iter() + .all(|(provider_id, managed)| provider_ids.contains(provider_id) + && managed_domain_ids.contains(&managed.id))); + } + other => panic!("expected AmbiguousManagedDomain, got {other:?}"), + } + + for (_, managed) in longest_zones { + let mut active: dns_managed_domains::ActiveModel = managed.into(); + active.auto_manage = Set(false); + active + .update(db.as_ref()) + .await + .expect("make tied longest zone ineligible"); + } + let (provider, managed) = service + .find_provider_for_domain("app.api.example.com") + .await + .expect("lookup with only parent eligible") + .expect("shorter parent must remain a valid fallback"); + assert_eq!(provider.name, "parent"); + assert_eq!(managed.domain, "example.com"); +} diff --git a/crates/temps-domains/Cargo.toml b/crates/temps-domains/Cargo.toml index 8e9d2b821..317a6e9b0 100644 --- a/crates/temps-domains/Cargo.toml +++ b/crates/temps-domains/Cargo.toml @@ -58,3 +58,6 @@ webpki-roots = "1.0" bollard = { workspace = true } futures-util = { workspace = true } tar = { workspace = true } +tower.workspace = true +wiremock = "0.6" +serial_test = "4.0" diff --git a/crates/temps-domains/src/dns_provider.rs b/crates/temps-domains/src/dns_provider.rs index 4f9d10330..08d039ae7 100644 --- a/crates/temps-domains/src/dns_provider.rs +++ b/crates/temps-domains/src/dns_provider.rs @@ -174,8 +174,8 @@ impl DnsProviderService for CloudflareDnsProvider { .join("."); info!( - "Setting TXT record for zone: {} base_domain: {} name: {} value: {}", - zone_id, base_domain, name, value + "Setting TXT record for zone: {} base_domain: {} name: {} value: [REDACTED]", + zone_id, base_domain, name ); // Get all existing TXT records with this name (try both full name and relative name) @@ -524,9 +524,7 @@ impl CloudflareDnsProvider { ); for record in &txt_records { - if let dns::DnsContent::TXT { content } = &record.content { - info!(" - TXT record id={} value={}", record.id, content); - } + info!(" - TXT record id={} value=[REDACTED]", record.id); } Ok(txt_records) diff --git a/crates/temps-domains/src/domain_service.rs b/crates/temps-domains/src/domain_service.rs index 3e4df33d3..652c9be73 100644 --- a/crates/temps-domains/src/domain_service.rs +++ b/crates/temps-domains/src/domain_service.rs @@ -325,7 +325,7 @@ impl DomainService { ); for (i, txt_record) in challenge_data.dns_txt_records.iter().enumerate() { - info!(" [{}] {} = {}", i + 1, txt_record.name, txt_record.value); + info!(" [{}] {} = [REDACTED]", i + 1, txt_record.name); } } } diff --git a/crates/temps-domains/src/handlers/domain_handler.rs b/crates/temps-domains/src/handlers/domain_handler.rs index bf2da815b..63d8b864c 100644 --- a/crates/temps-domains/src/handlers/domain_handler.rs +++ b/crates/temps-domains/src/handlers/domain_handler.rs @@ -1863,6 +1863,7 @@ async fn setup_dns_challenge( Json(request): Json, ) -> Result { permission_guard!(auth, DomainsWrite); + permission_guard!(auth, DnsProvidersWrite); // Check if DNS provider service is available let dns_provider_service = app_state.dns_provider_service.as_ref().ok_or_else(|| { @@ -1954,6 +1955,47 @@ async fn setup_dns_challenge( .build() })?; + // This governance check intentionally precedes zone lookup and provider-client + // construction, which decrypts credentials and can initiate external calls. + ensure_dns_provider_active(&dns_provider, &domain.domain, domain_id)?; + + let managed_domain = dns_provider_service + .find_verified_zone_for_provider(request.dns_provider_id, &domain.domain) + .await + .map_err(|e| { + error!( + "Failed to find a verified DNS zone for domain {} and provider {}: {}", + domain.domain, request.dns_provider_id, e + ); + match e { + temps_dns::errors::DnsError::AmbiguousManagedDomain { .. } => { + ErrorBuilder::new(StatusCode::CONFLICT) + .title("Ambiguous Managed DNS Zone") + .detail(format!( + "Multiple verified managed DNS zones match domain {} for provider {}. Remove the duplicate managed-domain entries and retry.", + domain.domain, request.dns_provider_id + )) + .build() + } + _ => ErrorBuilder::new(StatusCode::INTERNAL_SERVER_ERROR) + .title("DNS Zone Lookup Failed") + .detail(format!( + "Failed to verify that DNS provider {} manages domain {}", + request.dns_provider_id, domain.domain + )) + .build(), + } + })? + .ok_or_else(|| { + ErrorBuilder::new(StatusCode::BAD_REQUEST) + .title("DNS Provider Does Not Manage Domain") + .detail(format!( + "DNS provider {} has no verified zone covering domain {}", + request.dns_provider_id, domain.domain + )) + .build() + })?; + // Create DNS provider instance let provider_instance = dns_provider_service .create_provider_instance(&dns_provider) @@ -1965,8 +2007,7 @@ async fn setup_dns_challenge( .build() })?; - // Extract the base domain for the DNS provider - let base_domain = extract_base_domain(&domain.domain); + let authoritative_zone = managed_domain.domain; info!( "Setting up {} DNS TXT record(s) for {} using provider {}", @@ -1975,8 +2016,12 @@ async fn setup_dns_challenge( dns_provider.name ); - let (results, records_created) = - setup_dns_txt_records(provider_instance.as_ref(), &base_domain, &dns_txt_records).await; + let (results, records_created) = setup_dns_txt_records( + provider_instance.as_ref(), + &authoritative_zone, + &dns_txt_records, + ) + .await; let total_records = dns_txt_records.len() as u32; let all_success = records_created == total_records; @@ -2020,17 +2065,22 @@ async fn setup_dns_challenge( Ok(Json(response)) } -/// Extract the base domain from a full domain name (e.g., "sub.example.com" -> "example.com") -pub(crate) fn extract_base_domain(domain: &str) -> String { - // Handle wildcard domains - let domain = domain.strip_prefix("*.").unwrap_or(domain); - - let parts: Vec<&str> = domain.split('.').collect(); - if parts.len() >= 2 { - parts[parts.len() - 2..].join(".") - } else { - domain.to_string() +fn ensure_dns_provider_active( + provider: &temps_entities::dns_providers::Model, + domain: &str, + domain_id: i32, +) -> Result<(), Problem> { + if provider.is_active { + return Ok(()); } + + Err(ErrorBuilder::new(StatusCode::BAD_REQUEST) + .title("DNS Provider Is Inactive") + .detail(format!( + "DNS provider {} ({}) is inactive and cannot set up the DNS challenge for domain {} (ID {})", + provider.id, provider.name, domain, domain_id + )) + .build()) } /// Extract the record name relative to the base domain @@ -2089,6 +2139,82 @@ pub(crate) async fn setup_dns_txt_records( (results, records_created) } +pub(crate) enum DnsAutomationAuthorization { + Allowed, + Denied(String), + AuthorizationError(String), +} + +/// Validate and authorize an unattended DNS mutation before callers construct +/// a provider client. This function is deliberately provider-independent so a +/// denied or failed decision cannot decrypt provider credentials. +pub(crate) async fn authorize_dns_automation_request( + gate: &dyn temps_core::DnsAutomationGate, + request: &temps_core::DnsAutomationRequest, + actual_provider_id: i32, +) -> DnsAutomationAuthorization { + if let Err(reason) = validate_dns_automation_request(request, actual_provider_id) { + return DnsAutomationAuthorization::Denied(reason); + } + match gate.authorize(request).await { + Ok(temps_core::DnsAutomationDecision::Allow) => DnsAutomationAuthorization::Allowed, + // Policy implementations receive the ACME proof in `request`. Their + // free-form reason must never cross into logs or durable audit data, + // because a buggy implementation could reflect that proof verbatim. + Ok(temps_core::DnsAutomationDecision::Deny { .. }) => { + DnsAutomationAuthorization::Denied("automation policy denied the request".to_string()) + } + Err(_) => DnsAutomationAuthorization::AuthorizationError( + "automation policy evaluation failed".to_string(), + ), + } +} + +fn normalize_dns_name(name: &str) -> String { + name.trim().trim_end_matches('.').to_ascii_lowercase() +} + +pub(crate) fn validate_dns_automation_request( + request: &temps_core::DnsAutomationRequest, + actual_provider_id: i32, +) -> Result<(), String> { + if request.purpose != temps_core::DnsAutomationPurpose::AcmeDns01 { + return Err("background DNS mutation boundary accepts only ACME DNS-01 requests".into()); + } + if request.provider_id != actual_provider_id { + return Err("authorized DNS provider does not match the provider instance".into()); + } + if request.mutations.is_empty() { + return Err("background DNS mutation batch must not be empty".into()); + } + + let zone = normalize_dns_name(&request.zone); + let domain = normalize_dns_name(request.domain.trim_start_matches("*.")); + if zone.is_empty() + || domain.is_empty() + || (domain != zone && !domain.ends_with(&format!(".{zone}"))) + { + return Err("request domain is not covered by the authoritative DNS zone".into()); + } + + let expected_name = format!("_acme-challenge.{domain}"); + for mutation in &request.mutations { + if !mutation.record_type.eq_ignore_ascii_case("TXT") { + return Err("background DNS mutation boundary accepts only TXT records".into()); + } + if mutation.value.trim().is_empty() { + return Err("ACME DNS-01 mutation values must not be empty".into()); + } + if normalize_dns_name(&mutation.name) != expected_name { + return Err(format!( + "DNS mutation name must be exactly {expected_name} for this ACME authorization" + )); + } + } + + Ok(()) +} + /// Create a single ACME challenge TXT record using the DNS provider. /// Callers must remove stale records for every name in the batch before calling this /// (see `setup_dns_txt_records`) -- removing here, per-record, would delete a sibling @@ -2104,8 +2230,8 @@ async fn create_acme_txt_record( let record_name = acme_txt_record_name(base_domain, name); debug!( - "Creating TXT record: name={} (relative: {}), value={}, base_domain={}", - name, record_name, value, base_domain + "Creating TXT record: name={} (relative: {}), base_domain={}", + name, record_name, base_domain ); let request = DnsRecordRequest { @@ -2120,8 +2246,8 @@ async fn create_acme_txt_record( match provider.create_record(base_domain, request).await { Ok(_record) => { info!( - "Successfully created TXT record {} = {} for {}", - name, value, base_domain + "Successfully created TXT record {} for {}", + name, base_domain ); DnsChallengeRecordResult { name: name.to_string(), @@ -2366,6 +2492,38 @@ mod tests { }; use temps_dns::DnsError; + #[test] + fn inactive_dns_provider_is_rejected_with_context() { + let now = chrono::Utc::now(); + let provider = temps_entities::dns_providers::Model { + id: 42, + name: "disabled-cloudflare".to_string(), + provider_type: "cloudflare".to_string(), + credentials: "encrypted".to_string(), + is_active: false, + description: None, + last_used_at: None, + last_error: None, + created_at: now, + updated_at: now, + }; + + let problem = ensure_dns_provider_active(&provider, "api.example.com", 17) + .expect_err("inactive providers must be rejected before DNS setup"); + + assert_eq!(problem.status_code, StatusCode::BAD_REQUEST); + assert_eq!( + problem.body.get("title").and_then(|value| value.as_str()), + Some("DNS Provider Is Inactive") + ); + let detail = problem.body.get("detail").and_then(|value| value.as_str()); + assert!(detail.is_some_and(|detail| { + detail.contains("42 (disabled-cloudflare)") + && detail.contains("api.example.com") + && detail.contains("ID 17") + })); + } + /// In-memory DNS provider used to drive `setup_dns_txt_records` end-to-end without /// a live Cloudflare/Route53/etc. account. Mirrors the real providers' semantics: /// `list_records`/`remove_record` see every record in the zone, `create_record` @@ -2373,6 +2531,7 @@ mod tests { struct MockDnsProvider { records: Mutex>, next_id: AtomicU32, + fail_create_after: Option, } impl MockDnsProvider { @@ -2380,6 +2539,7 @@ mod tests { Self { records: Mutex::new(Vec::new()), next_id: AtomicU32::new(1), + fail_create_after: None, } } @@ -2389,6 +2549,14 @@ mod tests { provider } + fn failing_after(successful_creates: u32) -> Self { + Self { + records: Mutex::new(Vec::new()), + next_id: AtomicU32::new(1), + fail_create_after: Some(successful_creates), + } + } + fn record_names(&self) -> Vec<(String, String)> { self.records .lock() @@ -2461,6 +2629,12 @@ mod tests { request: DnsRecordRequest, ) -> Result { let id = self.next_id.fetch_add(1, Ordering::SeqCst).to_string(); + if self + .fail_create_after + .is_some_and(|limit| id.parse::().unwrap_or(u32::MAX) > limit) + { + return Err(DnsError::ApiError("injected create failure".to_string())); + } let record = DnsRecord { id: Some(id), zone: domain.to_string(), @@ -2574,4 +2748,232 @@ mod tests { .any(|(n, v)| n == "_acme-challenge" && v == "fresh-token")); assert!(remaining.iter().any(|(n, _)| n == "www")); } + + struct TestAutomationGate { + decision: Result, + } + + struct PanicAutomationGate; + + #[async_trait] + impl temps_core::DnsAutomationGate for PanicAutomationGate { + async fn authorize( + &self, + _request: &temps_core::DnsAutomationRequest, + ) -> Result { + panic!("invalid requests must be rejected before the authorization gate") + } + } + + #[async_trait] + impl temps_core::DnsAutomationGate for TestAutomationGate { + async fn authorize( + &self, + request: &temps_core::DnsAutomationRequest, + ) -> Result { + self.decision.clone().map_err(|reason| { + temps_core::DnsAutomationError::policy_evaluation_failed(request, reason) + }) + } + } + + fn automation_request(records: &[(String, String)]) -> temps_core::DnsAutomationRequest { + temps_core::DnsAutomationRequest { + purpose: temps_core::DnsAutomationPurpose::AcmeDns01, + domain: "*.example.com".to_string(), + zone: "example.com".to_string(), + provider_id: 7, + provider_name: "test".to_string(), + mutations: records + .iter() + .map(|(name, value)| temps_core::DnsAutomationMutation { + record_type: "TXT".to_string(), + name: name.clone(), + value: value.clone(), + }) + .collect(), + } + } + + #[tokio::test] + async fn denied_automation_does_not_touch_provider() { + let provider = MockDnsProvider::seed(vec![txt_record("1", "_acme-challenge", "stale")]); + let records = vec![( + "_acme-challenge.example.com".to_string(), + "fresh".to_string(), + )]; + let result = authorize_dns_automation_request( + &TestAutomationGate { + decision: Ok(temps_core::DnsAutomationDecision::Deny { + reason: "fresh".to_string(), + }), + }, + &automation_request(&records), + 7, + ) + .await; + + assert!(matches!( + result, + DnsAutomationAuthorization::Denied(reason) + if reason == "automation policy denied the request" && !reason.contains("fresh") + )); + assert_eq!( + provider.record_names(), + vec![("_acme-challenge".to_string(), "stale".to_string())] + ); + } + + #[tokio::test] + async fn authorization_error_does_not_touch_provider() { + let provider = MockDnsProvider::seed(vec![txt_record("1", "_acme-challenge", "stale")]); + let records = vec![( + "_acme-challenge.example.com".to_string(), + "fresh".to_string(), + )]; + let result = authorize_dns_automation_request( + &TestAutomationGate { + decision: Err("fresh".to_string()), + }, + &automation_request(&records), + 7, + ) + .await; + + assert!(matches!( + result, + DnsAutomationAuthorization::AuthorizationError(reason) + if reason == "automation policy evaluation failed" && !reason.contains("fresh") + )); + assert_eq!(provider.record_names().len(), 1); + } + + #[tokio::test] + async fn invalid_automation_requests_do_not_touch_gate_or_provider() { + let invalid_cases = [ + automation_request(&[]), + automation_request(&[("_acme-challenge.example.com".to_string(), " ".to_string())]), + automation_request(&[( + "_acme-challenge.attacker.example".to_string(), + "token".to_string(), + )]), + ]; + + for request in invalid_cases { + let provider = MockDnsProvider::seed(vec![txt_record("1", "_acme-challenge", "stale")]); + let result = authorize_dns_automation_request(&PanicAutomationGate, &request, 7).await; + + assert!(matches!(result, DnsAutomationAuthorization::Denied(_))); + assert_eq!( + provider.record_names(), + vec![("_acme-challenge".to_string(), "stale".to_string())] + ); + } + } + + #[tokio::test] + async fn provider_identity_mismatch_does_not_touch_gate_or_provider() { + let provider = MockDnsProvider::new(); + let request = automation_request(&[( + "_acme-challenge.example.com".to_string(), + "token".to_string(), + )]); + + let result = authorize_dns_automation_request(&PanicAutomationGate, &request, 99).await; + + assert!(matches!(result, DnsAutomationAuthorization::Denied(_))); + assert!(provider.record_names().is_empty()); + } + + #[tokio::test] + async fn allowed_automation_replaces_only_acme_txt_records() { + let provider = MockDnsProvider::seed(vec![ + txt_record("1", "_acme-challenge", "stale"), + txt_record("2", "www", "unrelated"), + ]); + let records = vec![( + "_acme-challenge.example.com".to_string(), + "fresh".to_string(), + )]; + let authorization = authorize_dns_automation_request( + &TestAutomationGate { + decision: Ok(temps_core::DnsAutomationDecision::Allow), + }, + &automation_request(&records), + 7, + ) + .await; + assert!(matches!(authorization, DnsAutomationAuthorization::Allowed)); + let (results, records_created) = + setup_dns_txt_records(&provider, "example.com", &records).await; + assert_eq!(records_created, 1); + assert!(results.iter().all(|result| result.success)); + let remaining = provider.record_names(); + assert!(remaining + .iter() + .any(|(name, value)| { name == "_acme-challenge" && value == "fresh" })); + assert!(remaining.iter().any(|(name, _)| name == "www")); + } + + #[tokio::test] + async fn authorized_setup_uses_the_request_authoritative_zone() { + let provider = MockDnsProvider::new(); + let mut request = automation_request(&[( + "_acme-challenge.api.dev.example.com".to_string(), + "fresh".to_string(), + )]); + request.domain = "api.dev.example.com".to_string(); + request.zone = "dev.example.com".to_string(); + + let authorization = authorize_dns_automation_request( + &TestAutomationGate { + decision: Ok(temps_core::DnsAutomationDecision::Allow), + }, + &request, + 7, + ) + .await; + + assert!(matches!(authorization, DnsAutomationAuthorization::Allowed)); + let records = request + .mutations + .iter() + .map(|mutation| (mutation.name.clone(), mutation.value.clone())) + .collect::>(); + setup_dns_txt_records(&provider, &request.zone, &records).await; + assert_eq!( + provider.record_names(), + vec![("_acme-challenge.api".to_string(), "fresh".to_string())] + ); + } + + #[tokio::test] + async fn partial_publish_failure_is_reported() { + let provider = MockDnsProvider::failing_after(1); + let records = vec![ + ( + "_acme-challenge.example.com".to_string(), + "first".to_string(), + ), + ( + "_acme-challenge.example.com".to_string(), + "second".to_string(), + ), + ]; + let request = automation_request(&records); + let authorization = authorize_dns_automation_request( + &TestAutomationGate { + decision: Ok(temps_core::DnsAutomationDecision::Allow), + }, + &request, + 7, + ) + .await; + assert!(matches!(authorization, DnsAutomationAuthorization::Allowed)); + let (results, records_created) = + setup_dns_txt_records(&provider, &request.zone, &records).await; + assert_eq!(records_created, 1); + assert_eq!(results.len(), 2); + assert!(!results[1].success); + } } diff --git a/crates/temps-domains/src/plugin.rs b/crates/temps-domains/src/plugin.rs index 8dd74279e..6e26bfa3b 100644 --- a/crates/temps-domains/src/plugin.rs +++ b/crates/temps-domains/src/plugin.rs @@ -70,6 +70,9 @@ impl TempsPlugin for DomainsPlugin { // wired into the TLS service so DNS-01 background renewals can auto-publish // the challenge TXT record when a DNS provider manages the domain's zone. let dns_provider_service = context.require_service::(); + let dns_automation_gate = + context.require_service::(); + let audit_service = context.require_service::(); // Create domain service first so the TLS service can drive the order-based // ACME flow during background HTTP-01 renewals (keeps auto-renewals @@ -93,7 +96,9 @@ impl TempsPlugin for DomainsPlugin { })? .with_domain_service(domain_service.clone()) .with_config_service(config_service.clone()) - .with_dns_provider_service(dns_provider_service.clone()); + .with_dns_provider_service(dns_provider_service.clone()) + .with_dns_automation_gate(dns_automation_gate) + .with_audit_logger(audit_service.clone()); // Add notification service if available if let Some(notif_service) = notification_service { @@ -112,8 +117,6 @@ impl TempsPlugin for DomainsPlugin { // The scheduler handles both initial check and daily scheduled checks // Get audit service - let audit_service = context.require_service::(); - // Get telemetry reporter (optional — default to noop so domains never hard-fail // if the telemetry plugin isn't registered) let telemetry = context diff --git a/crates/temps-domains/src/tls/service.rs b/crates/temps-domains/src/tls/service.rs index a88864d65..eee467553 100644 --- a/crates/temps-domains/src/tls/service.rs +++ b/crates/temps-domains/src/tls/service.rs @@ -6,10 +6,13 @@ use hickory_resolver::config::{ use hickory_resolver::net::runtime::TokioRuntimeProvider; use hickory_resolver::Resolver; use rustls::pki_types::{CertificateDer, PrivateKeyDer}; +use serde::Serialize; use std::sync::Arc; +use std::time::Duration; use temps_core::notifications::{ NotificationData, NotificationPriority, NotificationService, NotificationType, }; +use temps_core::{AuditLogger, AuditOperation}; use tracing::{error, info, warn}; use super::errors::{BuilderError, TlsError}; @@ -20,6 +23,61 @@ use super::repository::CertificateRepository; /// Type alias for the Tokio-based DNS resolver type TokioResolver = Resolver; +#[derive(Debug, Serialize)] +struct DnsAutomationAudit { + domain: String, + zone: String, + provider_id: i32, + provider_name: String, + outcome: String, + reason: Option, + #[serde(serialize_with = "serialize_redacted_mutations")] + mutations: Vec, +} + +fn serialize_redacted_mutations( + mutations: &[temps_core::DnsAutomationMutation], + serializer: S, +) -> Result +where + S: serde::Serializer, +{ + #[derive(Serialize)] + struct RedactedMutation<'a> { + record_type: &'a str, + name: &'a str, + value: &'static str, + } + + mutations + .iter() + .map(|mutation| RedactedMutation { + record_type: &mutation.record_type, + name: &mutation.name, + value: "[REDACTED]", + }) + .collect::>() + .serialize(serializer) +} + +impl AuditOperation for DnsAutomationAudit { + fn operation_type(&self) -> String { + "DNS_AUTOMATION_ACME_DNS01".to_string() + } + fn user_id(&self) -> Option { + None + } + fn ip_address(&self) -> Option { + None + } + fn user_agent(&self) -> &str { + "temps-certificate-renewal-scheduler" + } + fn serialize(&self) -> anyhow::Result { + serde_json::to_string(self).map_err(Into::into) + } +} + pub struct TlsService { repository: Arc, cert_provider: Arc, @@ -39,6 +97,11 @@ pub struct TlsService { /// provider manages the domain, DNS-01 renewals fall back to the manual /// "add this TXT record yourself" notification. dns_provider_service: Option>, + /// Fail-closed authorization gate for unattended DNS mutations. Human + /// provider management is governed independently by API permissions. + dns_automation_gate: Option>, + audit_logger: Option>, + dns_propagation_delay: Duration, } impl TlsService { @@ -74,6 +137,9 @@ impl TlsService { db: None, domain_service: None, dns_provider_service: None, + dns_automation_gate: None, + audit_logger: None, + dns_propagation_delay: Duration::from_secs(30), } } @@ -103,6 +169,32 @@ impl TlsService { self } + pub fn with_dns_automation_gate( + mut self, + dns_automation_gate: Arc, + ) -> Self { + self.dns_automation_gate = Some(dns_automation_gate); + self + } + + pub fn with_audit_logger(mut self, audit_logger: Arc) -> Self { + self.audit_logger = Some(audit_logger); + self + } + + async fn audit_dns_automation(&self, audit: DnsAutomationAudit) { + let Some(logger) = &self.audit_logger else { + warn!( + "DNS automation audit logger is unavailable; outcome={}", + audit.outcome + ); + return; + }; + if let Err(error) = logger.create_audit_log(&audit).await { + error!("Failed to persist DNS automation audit event: {error}"); + } + } + pub fn with_domain_service(mut self, domain_service: Arc) -> Self { self.domain_service = Some(domain_service); self @@ -655,7 +747,7 @@ impl TlsService { dns_provider_service: &Arc, report: &mut RenewalReport, ) -> bool { - let (provider, _managed_domain) = match dns_provider_service + let (provider, managed_domain) = match dns_provider_service .find_provider_for_domain(&cert.domain) .await { @@ -667,18 +759,7 @@ impl TlsService { } }; - let provider_instance = match dns_provider_service.create_provider_instance(&provider) { - Ok(instance) => instance, - Err(e) => { - warn!( - "Failed to initialize DNS provider {} for {}: {}", - provider.name, cert.domain, e - ); - return false; - } - }; - - let base_domain = crate::handlers::domain_handler::extract_base_domain(&cert.domain); + let authoritative_zone = managed_domain.domain; info!( "🔄 Auto-renewing DNS-01 certificate for {} via DNS provider {}", @@ -745,9 +826,97 @@ impl TlsService { .map(|record| (record.name.clone(), record.value.clone())) .collect(); + let mutations = dns_txt_records + .iter() + .map(|(name, value)| temps_core::DnsAutomationMutation { + record_type: "TXT".to_string(), + name: name.clone(), + value: value.clone(), + }) + .collect::>(); + let authorization_request = temps_core::DnsAutomationRequest { + purpose: temps_core::DnsAutomationPurpose::AcmeDns01, + domain: cert.domain.clone(), + zone: authoritative_zone.clone(), + provider_id: provider.id, + provider_name: provider.name.clone(), + mutations: mutations.clone(), + }; + let Some(gate) = &self.dns_automation_gate else { + self.audit_dns_automation(DnsAutomationAudit { + domain: cert.domain.clone(), + zone: authoritative_zone.clone(), + provider_id: provider.id, + provider_name: provider.name.clone(), + outcome: "denied".to_string(), + reason: Some("no automation gate is configured".to_string()), + mutations, + }) + .await; + return false; + }; + match crate::handlers::domain_handler::authorize_dns_automation_request( + gate.as_ref(), + &authorization_request, + provider.id, + ) + .await + { + crate::handlers::domain_handler::DnsAutomationAuthorization::Allowed => {} + crate::handlers::domain_handler::DnsAutomationAuthorization::Denied(reason) => { + self.audit_dns_automation(DnsAutomationAudit { + domain: cert.domain.clone(), + zone: authoritative_zone.clone(), + provider_id: provider.id, + provider_name: provider.name.clone(), + outcome: "denied".to_string(), + reason: Some(reason.clone()), + mutations: mutations.clone(), + }) + .await; + info!( + "Unattended DNS-01 renewal is not authorized for {} via provider {}: {}", + cert.domain, provider.name, reason + ); + return false; + } + crate::handlers::domain_handler::DnsAutomationAuthorization::AuthorizationError( + reason, + ) => { + self.audit_dns_automation(DnsAutomationAudit { + domain: cert.domain.clone(), + zone: authoritative_zone.clone(), + provider_id: provider.id, + provider_name: provider.name.clone(), + outcome: "authorization_error".to_string(), + reason: Some(reason.clone()), + mutations: mutations.clone(), + }) + .await; + warn!( + "Failed to authorize unattended DNS-01 renewal for {} via provider {}: {}", + cert.domain, provider.name, reason + ); + return false; + } + } + + // Provider construction decrypts credentials, so it must happen only + // after the unattended mutation policy has explicitly allowed this + // exact provider, zone, domain, and record batch. + let provider_instance = match dns_provider_service.create_provider_instance(&provider) { + Ok(instance) => instance, + Err(error) => { + warn!( + "Failed to initialize DNS provider {} for {} after automation authorization: {}", + provider.name, cert.domain, error + ); + return false; + } + }; let (results, records_created) = crate::handlers::domain_handler::setup_dns_txt_records( provider_instance.as_ref(), - &base_domain, + &authoritative_zone, &dns_txt_records, ) .await; @@ -781,15 +950,36 @@ impl TlsService { &cert.verification_method, ) .await; + self.audit_dns_automation(DnsAutomationAudit { + domain: cert.domain.clone(), + zone: authoritative_zone.clone(), + provider_id: provider.id, + provider_name: provider.name.clone(), + outcome: "publish_failed".to_string(), + reason: Some(error_msg), + mutations: authorization_request.mutations.clone(), + }) + .await; return true; } + self.audit_dns_automation(DnsAutomationAudit { + domain: cert.domain.clone(), + zone: authoritative_zone.clone(), + provider_id: provider.id, + provider_name: provider.name.clone(), + outcome: "published".to_string(), + reason: None, + mutations: authorization_request.mutations.clone(), + }) + .await; + // Step 3: Give DNS a moment to propagate before asking Let's Encrypt to validate. info!( "Waiting for DNS propagation before validating DNS-01 challenge for {}...", cert.domain ); - tokio::time::sleep(tokio::time::Duration::from_secs(30)).await; + tokio::time::sleep(self.dns_propagation_delay).await; // Step 4: Accept the challenge and finalize the persisted order. match domain_service @@ -1324,6 +1514,693 @@ fn load_private_key(content: &[u8]) -> Result, TlsError> #[cfg(test)] mod tests { use super::*; + use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering}; + use std::sync::Mutex; + + #[test] + fn dns_automation_audit_redacts_acme_values() { + let audit = DnsAutomationAudit { + domain: "example.com".to_string(), + zone: "example.com".to_string(), + provider_id: 7, + provider_name: "production-dns".to_string(), + outcome: "published".to_string(), + reason: None, + mutations: vec![temps_core::DnsAutomationMutation { + record_type: "TXT".to_string(), + name: "_acme-challenge.example.com".to_string(), + value: "super-secret-acme-token".to_string(), + }], + }; + + let serialized = AuditOperation::serialize(&audit).unwrap(); + + assert!(!serialized.contains("super-secret-acme-token")); + assert!(serialized.contains("[REDACTED]")); + assert!(serialized.contains("_acme-challenge.example.com")); + } + + #[derive(Default)] + struct DenyingDnsAutomationGate { + requests: Mutex>, + } + + #[async_trait::async_trait] + impl temps_core::DnsAutomationGate for DenyingDnsAutomationGate { + async fn authorize( + &self, + request: &temps_core::DnsAutomationRequest, + ) -> Result { + self.requests.lock().unwrap().push(request.clone()); + Ok(temps_core::DnsAutomationDecision::Deny { + reason: "scheduler principal lacks dns:automation:write".to_string(), + }) + } + } + + struct ErroringDnsAutomationGate; + + #[async_trait::async_trait] + impl temps_core::DnsAutomationGate for ErroringDnsAutomationGate { + async fn authorize( + &self, + request: &temps_core::DnsAutomationRequest, + ) -> Result { + Err(temps_core::DnsAutomationError::policy_evaluation_failed( + request, + "policy store offline", + )) + } + } + + struct AllowingDnsAutomationGate; + + #[async_trait::async_trait] + impl temps_core::DnsAutomationGate for AllowingDnsAutomationGate { + async fn authorize( + &self, + _request: &temps_core::DnsAutomationRequest, + ) -> Result { + Ok(temps_core::DnsAutomationDecision::Allow) + } + } + + #[derive(Default)] + struct RecordingAuditLogger { + operations: Mutex>, + } + + #[async_trait::async_trait] + impl temps_core::AuditLogger for RecordingAuditLogger { + async fn create_audit_log( + &self, + operation: &dyn temps_core::AuditOperation, + ) -> anyhow::Result<()> { + self.operations + .lock() + .unwrap() + .push((operation.operation_type(), operation.serialize()?)); + Ok(()) + } + } + + struct DnsRenewalCertificateProvider { + completion_calls: AtomicUsize, + } + + fn dns_renewal_test_server_config(data_dir: std::path::PathBuf) -> temps_config::ServerConfig { + temps_config::ServerConfig { + address: "127.0.0.1:0".to_string(), + database_url: "postgres://unused".to_string(), + tls_address: None, + console_address: "127.0.0.1:0".to_string(), + console_admin_address: None, + admin_allowed_ips: vec![], + admin_allowed_hosts: vec![], + admin_trust_forwarded_for: false, + data_dir, + auth_secret: "test-secret".to_string(), + encryption_key: "test-key".to_string(), + api_base_url: "/api".to_string(), + postgres_max_connections: None, + postgres_min_connections: None, + postgres_connect_timeout_secs: None, + postgres_acquire_timeout_secs: None, + postgres_idle_timeout_secs: None, + postgres_max_lifetime_secs: None, + clickhouse_url: None, + clickhouse_database: None, + clickhouse_user: None, + clickhouse_password: None, + docker_extra_networks: vec![], + } + } + + #[async_trait::async_trait] + impl CertificateProvider for DnsRenewalCertificateProvider { + async fn provision( + &self, + domain: &str, + challenge: ChallengeType, + _email: &str, + ) -> Result { + assert_eq!(challenge, ChallengeType::Dns01); + Ok(ProvisioningResult::Challenge(ChallengeData { + challenge_type: ChallengeType::Dns01, + domain: domain.to_string(), + token: "order-token".to_string(), + key_authorization: "key-authorization".to_string(), + validation_url: Some("https://acme.test/challenge/1".to_string()), + dns_txt_records: vec![crate::tls::models::DnsTxtRecord { + name: format!("_acme-challenge.{domain}"), + value: "secret-acme-proof".to_string(), + validation_url: "https://acme.test/challenge/1".to_string(), + }], + order_url: Some("https://acme.test/order/1".to_string()), + })) + } + + async fn complete_challenge( + &self, + domain: &str, + _challenge_data: &ChallengeData, + _email: &str, + ) -> Result { + self.completion_calls.fetch_add(1, AtomicOrdering::SeqCst); + Ok(Certificate { + id: 1, + domain: domain.to_string(), + certificate_pem: "certificate".to_string(), + private_key_pem: "private-key".to_string(), + expiration_time: chrono::Utc::now() + chrono::Duration::days(90), + last_renewed: Some(chrono::Utc::now()), + is_wildcard: false, + verification_method: "dns-01".to_string(), + status: CertificateStatus::Active, + }) + } + + fn supported_challenges(&self) -> Vec { + vec![ChallengeType::Dns01] + } + + async fn validate_prerequisites( + &self, + _domain: &str, + _email: &str, + ) -> Result { + Ok(ValidationResult { + is_valid: true, + errors: vec![], + warnings: vec![], + }) + } + + async fn cancel_order(&self, _domain: &str) -> Result<(), ProviderError> { + Ok(()) + } + } + + #[tokio::test] + #[serial_test::serial(dns_renewal_db)] + async fn test_dns01_renewal_policy_failures_precede_provider_credential_decryption() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_core::AppSettings; + use temps_entities::{dns_managed_domains, dns_providers, settings}; + + let test_db = match TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable( + &error.to_string(), + ) => + { + eprintln!("Docker unavailable; skipping DNS policy ordering test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password( + "dns-policy-ordering-test", + )); + let repository = Arc::new(MockCertificateRepository::new()); + let certificate_provider = Arc::new(DnsRenewalCertificateProvider { + completion_calls: AtomicUsize::new(0), + }); + let domain_service = Arc::new(crate::DomainService::new( + db.clone(), + certificate_provider.clone(), + repository.clone(), + encryption.clone(), + )); + let dns_provider_service = Arc::new(temps_dns::services::DnsProviderService::new( + db.clone(), + encryption.clone(), + )); + let mut app_settings = AppSettings::default(); + app_settings.letsencrypt.email = Some("acme@example.com".to_string()); + settings::ActiveModel { + id: Set(1), + data: Set(serde_json::to_value(app_settings).unwrap()), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert ACME settings"); + let config_dir = std::env::temp_dir().join(format!( + "temps-dns-policy-ordering-config-{}", + uuid::Uuid::new_v4() + )); + let config_service = Arc::new(temps_config::ConfigService::new( + Arc::new(dns_renewal_test_server_config(config_dir)), + db.clone(), + )); + + struct PolicyCase { + label: &'static str, + gate: Option>, + expected_reason: &'static str, + } + let cases = [ + PolicyCase { + label: "missing", + gate: None, + expected_reason: "no automation gate is configured", + }, + PolicyCase { + label: "denied", + gate: Some(Arc::new(DenyingDnsAutomationGate::default())), + expected_reason: "automation policy denied the request", + }, + PolicyCase { + label: "error", + gate: Some(Arc::new(ErroringDnsAutomationGate)), + expected_reason: "automation policy evaluation failed", + }, + ]; + + for PolicyCase { + label, + gate, + expected_reason, + } in cases + { + let zone = format!("{label}.example.com"); + let domain = format!("app.{zone}"); + let persisted_domain = domain_service + .create_domain(&domain, "dns-01") + .await + .expect("create DNS-01 domain for policy-ordering case"); + let provider = dns_providers::ActiveModel { + name: Set(format!("{label}-provider")), + provider_type: Set("cloudflare".to_string()), + // A policy-ordering regression tries to decrypt this and fails + // before it can produce the expected policy/manual outcome. + credentials: Set("not-valid-ciphertext".to_string()), + is_active: Set(true), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert active provider"); + dns_managed_domains::ActiveModel { + provider_id: Set(provider.id), + domain: Set(zone), + auto_manage: Set(true), + verified: Set(true), + generated_hostname_mode: Set("standard".to_string()), + sync_generated_records: Set(false), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert eligible managed zone"); + + let audit = Arc::new(RecordingAuditLogger::default()); + let mut service = TlsService::new(repository.clone(), certificate_provider.clone()) + .with_config_service(config_service.clone()) + .with_domain_service(domain_service.clone()) + .with_dns_provider_service(dns_provider_service.clone()) + .with_audit_logger(audit.clone()); + if let Some(gate) = gate { + service = service.with_dns_automation_gate(gate); + } + let certificate = Certificate { + id: persisted_domain.id, + domain: domain.clone(), + certificate_pem: "old-certificate".to_string(), + private_key_pem: "old-private-key".to_string(), + expiration_time: chrono::Utc::now() + chrono::Duration::days(7), + last_renewed: None, + is_wildcard: false, + verification_method: "dns-01".to_string(), + status: CertificateStatus::Active, + }; + let mut report = RenewalReport { + total_checked: 1, + auto_renewed: vec![], + renewal_failed: vec![], + manual_action_needed: vec![], + }; + + service + .handle_dns01_notification(&certificate, &mut report) + .await; + + assert_eq!(report.manual_action_needed.len(), 1, "case {label}"); + assert_eq!( + report.manual_action_needed[0].domain, domain, + "case {label}" + ); + let operations = audit.operations.lock().unwrap(); + assert_eq!(operations.len(), 1, "case {label}"); + assert!(operations[0].1.contains(expected_reason), "case {label}"); + } + } + + #[tokio::test] + #[serial_test::serial(dns_renewal_db)] + async fn test_check_and_renew_certificates_dns01_denied_gate_falls_back_to_manual_action() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_core::AppSettings; + use temps_dns::providers::{CloudflareCredentials, DnsProviderType, ProviderCredentials}; + use temps_dns::services::{ + AddManagedDomainRequest, CreateProviderRequest, DnsProviderService, + }; + use temps_entities::{dns_managed_domains, settings}; + + let test_db = match TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable( + &error.to_string(), + ) => + { + eprintln!("Docker unavailable; skipping DNS renewal regression test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password( + "dns-renewal-test", + )); + let repository = Arc::new(DefaultCertificateRepository::new( + db.clone(), + encryption.clone(), + )); + let certificate_provider = Arc::new(DnsRenewalCertificateProvider { + completion_calls: AtomicUsize::new(0), + }); + let domain_service = Arc::new(crate::DomainService::new( + db.clone(), + certificate_provider.clone(), + repository.clone(), + encryption.clone(), + )); + let domain = domain_service + .create_domain("app.example.com", "dns-01") + .await + .expect("create renewal domain"); + + let mut app_settings = AppSettings::default(); + app_settings.letsencrypt.email = Some("acme@example.com".to_string()); + settings::ActiveModel { + id: Set(1), + data: Set(serde_json::to_value(app_settings).unwrap()), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert ACME settings"); + + let config_dir = + std::env::temp_dir().join(format!("temps-dns-renewal-config-{}", uuid::Uuid::new_v4())); + let server_config = dns_renewal_test_server_config(config_dir.clone()); + let config_service = Arc::new(temps_config::ConfigService::new( + Arc::new(server_config), + db.clone(), + )); + + let dns_provider_service = + Arc::new(DnsProviderService::new(db.clone(), encryption.clone())); + let provider = dns_provider_service + .create(CreateProviderRequest { + name: "production-dns".to_string(), + provider_type: DnsProviderType::Cloudflare, + credentials: ProviderCredentials::Cloudflare(CloudflareCredentials { + api_token: "test-token".to_string(), + account_id: None, + }), + description: None, + }) + .await + .expect("create provider"); + let managed = dns_provider_service + .add_managed_domain( + provider.id, + AddManagedDomainRequest { + domain: "example.com".to_string(), + auto_manage: true, + generated_hostname_mode: None, + sync_generated_records: false, + }, + ) + .await + .expect("add managed zone"); + let mut managed_active: dns_managed_domains::ActiveModel = managed.into(); + managed_active.verified = Set(true); + managed_active + .update(db.as_ref()) + .await + .expect("verify managed zone"); + + let gate = Arc::new(DenyingDnsAutomationGate::default()); + let audit = Arc::new(RecordingAuditLogger::default()); + let service = TlsService::new(repository.clone(), certificate_provider.clone()) + .with_config_service(config_service) + .with_domain_service(domain_service.clone()) + .with_dns_provider_service(dns_provider_service.clone()) + .with_dns_automation_gate(gate.clone()) + .with_audit_logger(audit.clone()); + let certificate = Certificate { + id: domain.id, + domain: domain.domain, + certificate_pem: "old-certificate".to_string(), + private_key_pem: "old-private-key".to_string(), + expiration_time: chrono::Utc::now() + chrono::Duration::days(7), + last_renewed: None, + is_wildcard: false, + verification_method: "dns-01".to_string(), + status: CertificateStatus::Active, + }; + repository + .save_certificate(certificate) + .await + .expect("persist expiring DNS-01 certificate"); + + let report = service + .check_and_renew_certificates(30) + .await + .expect("run scheduled certificate renewal"); + + assert_eq!(report.total_checked, 1); + assert!(report.auto_renewed.is_empty()); + assert!(report.renewal_failed.is_empty()); + assert_eq!(report.manual_action_needed.len(), 1); + assert_eq!(report.manual_action_needed[0].domain, "app.example.com"); + assert_eq!( + certificate_provider + .completion_calls + .load(AtomicOrdering::SeqCst), + 0 + ); + let pending_order = repository + .find_acme_order_by_domain(domain.id) + .await + .expect("query pending order") + .expect("challenge request must remain recoverable"); + assert_eq!(pending_order.status, "pending"); + + let requests = gate.requests.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].domain, "app.example.com"); + assert_eq!(requests[0].zone, "example.com"); + assert_eq!(requests[0].provider_id, provider.id); + assert_eq!(requests[0].mutations[0].value, "secret-acme-proof"); + drop(requests); + + let operations = audit.operations.lock().unwrap(); + assert_eq!(operations.len(), 1); + assert_eq!(operations[0].0, "DNS_AUTOMATION_ACME_DNS01"); + assert!(operations[0].1.contains("\"outcome\":\"denied\"")); + assert!(operations[0] + .1 + .contains("automation policy denied the request")); + assert!(operations[0].1.contains("[REDACTED]")); + assert!(!operations[0].1.contains("secret-acme-proof")); + + let _ = std::fs::remove_dir_all(config_dir); + } + + #[tokio::test] + #[serial_test::serial(dns_renewal_db)] + async fn test_try_dns01_renewal_with_provider_publishes_and_finalizes_certificate() { + use sea_orm::{ActiveModelTrait, ActiveValue::Set}; + use temps_core::AppSettings; + use temps_dns::providers::{DnsProviderType, PebbleCredentials, ProviderCredentials}; + use temps_dns::services::{ + AddManagedDomainRequest, CreateProviderRequest, DnsProviderService, + }; + use temps_entities::{dns_managed_domains, settings}; + use wiremock::matchers::{body_json, method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let test_db = match TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable( + &error.to_string(), + ) => + { + eprintln!("Docker unavailable; skipping DNS renewal regression test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let dns_server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/clear-txt")) + .and(body_json(serde_json::json!({ + "host": "_acme-challenge.app.example.com." + }))) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&dns_server) + .await; + Mock::given(method("POST")) + .and(path("/set-txt")) + .and(body_json(serde_json::json!({ + "host": "_acme-challenge.app.example.com.", + "value": "secret-acme-proof" + }))) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&dns_server) + .await; + + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password( + "dns-renewal-success-test", + )); + let repository = Arc::new(DefaultCertificateRepository::new( + db.clone(), + encryption.clone(), + )); + let certificate_provider = Arc::new(DnsRenewalCertificateProvider { + completion_calls: AtomicUsize::new(0), + }); + let domain_service = Arc::new(crate::DomainService::new( + db.clone(), + certificate_provider.clone(), + repository.clone(), + encryption.clone(), + )); + let domain = domain_service + .create_domain("app.example.com", "dns-01") + .await + .expect("create renewal domain"); + + let mut app_settings = AppSettings::default(); + app_settings.letsencrypt.email = Some("acme@example.com".to_string()); + settings::ActiveModel { + id: Set(1), + data: Set(serde_json::to_value(app_settings).unwrap()), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert ACME settings"); + let config_dir = std::env::temp_dir().join(format!( + "temps-dns-renewal-success-config-{}", + uuid::Uuid::new_v4() + )); + let server_config = dns_renewal_test_server_config(config_dir.clone()); + let config_service = Arc::new(temps_config::ConfigService::new( + Arc::new(server_config), + db.clone(), + )); + + let dns_provider_service = + Arc::new(DnsProviderService::new(db.clone(), encryption.clone())); + let provider = dns_provider_service + .create(CreateProviderRequest { + name: "pebble-dns".to_string(), + provider_type: DnsProviderType::Pebble, + credentials: ProviderCredentials::Pebble(PebbleCredentials { + management_url: dns_server.uri(), + }), + description: None, + }) + .await + .expect("create provider"); + let managed = dns_provider_service + .add_managed_domain( + provider.id, + AddManagedDomainRequest { + domain: "example.com".to_string(), + auto_manage: true, + generated_hostname_mode: None, + sync_generated_records: false, + }, + ) + .await + .expect("add managed zone"); + let mut managed_active: dns_managed_domains::ActiveModel = managed.into(); + managed_active.verified = Set(true); + managed_active + .update(db.as_ref()) + .await + .expect("verify managed zone"); + + let audit = Arc::new(RecordingAuditLogger::default()); + let mut service = TlsService::new(repository, certificate_provider.clone()) + .with_config_service(config_service) + .with_dns_automation_gate(Arc::new(AllowingDnsAutomationGate)) + .with_audit_logger(audit.clone()); + service.dns_propagation_delay = tokio::time::Duration::ZERO; + let certificate = Certificate { + id: domain.id, + domain: domain.domain, + certificate_pem: "old-certificate".to_string(), + private_key_pem: "old-private-key".to_string(), + expiration_time: chrono::Utc::now() + chrono::Duration::days(7), + last_renewed: None, + is_wildcard: false, + verification_method: "dns-01".to_string(), + status: CertificateStatus::Active, + }; + std::env::set_var("TEMPS_ALLOW_PEBBLE_PROVIDER", "1"); + let task = tokio::spawn(async move { + let mut report = RenewalReport { + total_checked: 1, + auto_renewed: vec![], + renewal_failed: vec![], + manual_action_needed: vec![], + }; + let handled = service + .try_dns01_renewal_with_provider( + &certificate, + &domain_service, + &dns_provider_service, + &mut report, + ) + .await; + (handled, report) + }); + let (handled, report) = task.await.expect("renewal task"); + std::env::remove_var("TEMPS_ALLOW_PEBBLE_PROVIDER"); + + assert!(handled); + assert_eq!(report.auto_renewed, vec!["app.example.com"]); + assert!(report.renewal_failed.is_empty()); + assert!(report.manual_action_needed.is_empty()); + assert_eq!( + certificate_provider + .completion_calls + .load(AtomicOrdering::SeqCst), + 1 + ); + let operations = audit.operations.lock().unwrap(); + assert_eq!(operations.len(), 1); + assert!(operations[0].1.contains("\"outcome\":\"published\"")); + assert!(operations[0].1.contains("[REDACTED]")); + assert!(!operations[0].1.contains("secret-acme-proof")); + + let _ = std::fs::remove_dir_all(config_dir); + } use crate::tls::errors::ProviderError; use crate::tls::models::{ Certificate, CertificateFilter, CertificateStatus, ChallengeData, ChallengeType, diff --git a/crates/temps-domains/tests/dns_governance_router_test.rs b/crates/temps-domains/tests/dns_governance_router_test.rs new file mode 100644 index 000000000..7ce685a91 --- /dev/null +++ b/crates/temps-domains/tests/dns_governance_router_test.rs @@ -0,0 +1,437 @@ +use std::sync::{Arc, Mutex}; + +use async_trait::async_trait; +use axum::{ + body::Body, + http::{Method, Request, StatusCode}, +}; +use http_body_util::BodyExt; +use sea_orm::{ActiveModelTrait, ActiveValue::Set, DatabaseBackend, MockDatabase}; +use temps_auth::{AuthContext, Permission}; +use temps_core::{AuditLogger, AuditOperation, NoopTelemetryReporter, RequestMetadata}; +use temps_domains::tls::{ + ChallengeData, ChallengeType, ProviderError, ProvisioningResult, ValidationResult, +}; +use temps_domains::{ + Certificate, CertificateProvider, CertificateRepository, DefaultCertificateRepository, + DomainAppState, DomainService, TlsService, +}; +use tower::ServiceExt; + +struct UnusedCertificateProvider; + +#[async_trait] +impl CertificateProvider for UnusedCertificateProvider { + async fn provision( + &self, + _domain: &str, + _challenge: ChallengeType, + _email: &str, + ) -> Result { + panic!("permission/inactive-provider tests must not provision a certificate") + } + + async fn complete_challenge( + &self, + _domain: &str, + _challenge_data: &ChallengeData, + _email: &str, + ) -> Result { + panic!("permission/inactive-provider tests must not finalize a certificate") + } + + fn supported_challenges(&self) -> Vec { + vec![ChallengeType::Dns01] + } + + async fn validate_prerequisites( + &self, + _domain: &str, + _email: &str, + ) -> Result { + Ok(ValidationResult { + is_valid: true, + errors: vec![], + warnings: vec![], + }) + } + + async fn cancel_order(&self, _domain: &str) -> Result<(), ProviderError> { + Ok(()) + } +} + +#[derive(Default)] +struct RecordingAudit { + operations: Mutex>, +} + +#[async_trait] +impl AuditLogger for RecordingAudit { + async fn create_audit_log(&self, operation: &dyn AuditOperation) -> anyhow::Result<()> { + self.operations + .lock() + .unwrap() + .push(operation.operation_type()); + Ok(()) + } +} + +fn test_user() -> temps_entities::users::Model { + let now = chrono::Utc::now(); + temps_entities::users::Model { + id: 42, + name: "Domain operator".to_string(), + email: "domains@example.com".to_string(), + password_hash: None, + email_verified: true, + email_verification_token: None, + email_verification_expires: None, + password_reset_token: None, + password_reset_expires: None, + must_change_password: false, + deleted_at: None, + mfa_secret: None, + mfa_enabled: false, + mfa_recovery_codes: None, + oidc_subject: None, + oidc_provider_id: None, + created_at: now, + updated_at: now, + } +} + +fn auth(permissions: Vec) -> AuthContext { + AuthContext::new_api_key( + test_user(), + None, + Some(permissions), + "domain-governance-test".to_string(), + 1, + ) +} + +fn metadata() -> RequestMetadata { + RequestMetadata { + ip_address: "127.0.0.1".to_string(), + user_agent: "domain-governance-router-test".to_string(), + headers: Default::default(), + visitor_id_cookie: None, + session_id_cookie: None, + base_url: "http://localhost".to_string(), + scheme: "http".to_string(), + host: "localhost".to_string(), + is_secure: false, + } +} + +fn permission_router() -> axum::Router { + // Empty MockDatabase: if either guard is moved below service access, the + // test fails on an unexpected query instead of producing a false-positive 403. + let db = Arc::new(MockDatabase::new(DatabaseBackend::Postgres).into_connection()); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let repository = Arc::new(DefaultCertificateRepository::new( + db.clone(), + encryption.clone(), + )); + let provider = Arc::new(UnusedCertificateProvider); + let domain_service = Arc::new(DomainService::new( + db.clone(), + provider.clone(), + repository.clone(), + encryption, + )); + let state = Arc::new(DomainAppState { + tls_service: Arc::new(TlsService::new(repository.clone(), provider)), + repository, + domain_service, + dns_provider_service: None, + audit_service: Arc::new(RecordingAudit::default()), + telemetry: Arc::new(NoopTelemetryReporter), + }); + temps_domains::configure_routes().with_state(state) +} + +async fn setup_dns_status(router: axum::Router, permissions: Vec) -> StatusCode { + let mut request = Request::builder() + .method(Method::POST) + .uri("/domains/123/setup-dns") + .header("content-type", "application/json") + .body(Body::from(r#"{"dns_provider_id":7}"#)) + .unwrap(); + request.extensions_mut().insert(auth(permissions)); + request.extensions_mut().insert(metadata()); + router.oneshot(request).await.unwrap().status() +} + +#[tokio::test] +async fn test_setup_dns_with_only_domains_write_returns_forbidden_before_state_touch() { + let status = setup_dns_status(permission_router(), vec![Permission::DomainsWrite]).await; + + assert_eq!(status, StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn test_setup_dns_with_only_dns_providers_write_returns_forbidden_before_state_touch() { + let status = setup_dns_status(permission_router(), vec![Permission::DnsProvidersWrite]).await; + + assert_eq!(status, StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn test_setup_dns_inactive_provider_rejects_without_dns_or_audit_touch() { + use temps_domains::tls::models::AcmeOrder; + use temps_entities::{dns_managed_domains, dns_providers}; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping inactive-provider router test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let repository = Arc::new(DefaultCertificateRepository::new( + db.clone(), + encryption.clone(), + )); + let certificate_provider = Arc::new(UnusedCertificateProvider); + let domain_service = Arc::new(DomainService::new( + db.clone(), + certificate_provider.clone(), + repository.clone(), + encryption.clone(), + )); + let domain = domain_service + .create_domain("app.example.com", "dns-01") + .await + .expect("create domain"); + repository + .save_acme_order(AcmeOrder { + id: 0, + order_url: "https://acme.test/order/1".to_string(), + domain_id: domain.id, + email: "acme@example.com".to_string(), + status: "pending".to_string(), + identifiers: serde_json::json!([]), + authorizations: Some(serde_json::json!({ + "dns_txt_records": [{ + "name": "_acme-challenge.app.example.com", + "value": "must-not-be-published" + }] + })), + finalize_url: None, + certificate_url: None, + error: None, + error_type: None, + token: Some("token".to_string()), + key_authorization: Some("key-auth".to_string()), + created_at: chrono::Utc::now(), + updated_at: chrono::Utc::now(), + expires_at: Some(chrono::Utc::now() + chrono::Duration::days(1)), + }) + .await + .expect("save order"); + let provider = dns_providers::ActiveModel { + name: Set("disabled-provider".to_string()), + provider_type: Set("manual".to_string()), + // Deliberately invalid ciphertext: a regression that initializes the + // provider before checking is_active turns this request into a 500. + credentials: Set("not-valid-ciphertext".to_string()), + is_active: Set(false), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert inactive provider"); + dns_managed_domains::ActiveModel { + provider_id: Set(provider.id), + domain: Set("example.com".to_string()), + auto_manage: Set(false), + verified: Set(true), + generated_hostname_mode: Set("standard".to_string()), + sync_generated_records: Set(false), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert managed domain"); + + let dns_provider_service = Arc::new(temps_dns::services::DnsProviderService::new( + db.clone(), + encryption, + )); + let audit = Arc::new(RecordingAudit::default()); + let state = Arc::new(DomainAppState { + tls_service: Arc::new(TlsService::new(repository.clone(), certificate_provider)), + repository, + domain_service, + dns_provider_service: Some(dns_provider_service), + audit_service: audit.clone(), + telemetry: Arc::new(NoopTelemetryReporter), + }); + + let mut request = Request::builder() + .method(Method::POST) + .uri(format!("/domains/{}/setup-dns", domain.id)) + .header("content-type", "application/json") + .body(Body::from(format!( + r#"{{"dns_provider_id":{}}}"#, + provider.id + ))) + .unwrap(); + request.extensions_mut().insert(auth(vec![ + Permission::DomainsWrite, + Permission::DnsProvidersWrite, + ])); + request.extensions_mut().insert(metadata()); + let response = temps_domains::configure_routes() + .with_state(state) + .oneshot(request) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let body = response.into_body().collect().await.unwrap().to_bytes(); + let problem: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!(problem["title"], "DNS Provider Is Inactive"); + assert!(problem["detail"] + .as_str() + .unwrap() + .contains("disabled-provider")); + assert!( + audit.operations.lock().unwrap().is_empty(), + "rejected inactive providers must not emit a successful DNS setup audit" + ); +} + +#[tokio::test] +#[serial_test::serial(dns_governance_db)] +async fn test_setup_dns_normalized_zone_ambiguity_returns_conflict() { + use temps_domains::tls::models::AcmeOrder; + use temps_entities::{dns_managed_domains, dns_providers}; + + let test_db = match temps_database::test_utils::TestDatabase::with_migrations().await { + Ok(db) => db, + Err(error) + if temps_database::test_utils::is_container_runtime_unavailable(&error.to_string()) => + { + eprintln!("Docker unavailable; skipping setup-DNS ambiguity test: {error}"); + return; + } + Err(error) => panic!("failed to create test database: {error}"), + }; + let db = test_db.db.clone(); + let encryption = Arc::new(temps_core::EncryptionService::new_from_password("test")); + let repository = Arc::new(DefaultCertificateRepository::new( + db.clone(), + encryption.clone(), + )); + let certificate_provider = Arc::new(UnusedCertificateProvider); + let domain_service = Arc::new(DomainService::new( + db.clone(), + certificate_provider.clone(), + repository.clone(), + encryption.clone(), + )); + let domain = domain_service + .create_domain("app.example.com", "dns-01") + .await + .expect("create domain"); + repository + .save_acme_order(AcmeOrder { + id: 0, + order_url: "https://acme.test/order/ambiguity".to_string(), + domain_id: domain.id, + email: "acme@example.com".to_string(), + status: "pending".to_string(), + identifiers: serde_json::json!([]), + authorizations: Some(serde_json::json!({ + "dns_txt_records": [{ + "name": "_acme-challenge.app.example.com", + "value": "must-not-be-published" + }] + })), + finalize_url: None, + certificate_url: None, + error: None, + error_type: None, + token: Some("token".to_string()), + key_authorization: Some("key-auth".to_string()), + created_at: chrono::Utc::now(), + updated_at: chrono::Utc::now(), + expires_at: Some(chrono::Utc::now() + chrono::Duration::days(1)), + }) + .await + .expect("save order"); + let provider = dns_providers::ActiveModel { + name: Set("ambiguous-provider".to_string()), + provider_type: Set("manual".to_string()), + credentials: Set("not-valid-ciphertext".to_string()), + is_active: Set(true), + description: Set(None), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert provider"); + for zone in ["example.com", " *.EXAMPLE.COM. "] { + dns_managed_domains::ActiveModel { + provider_id: Set(provider.id), + domain: Set(zone.to_string()), + auto_manage: Set(false), + verified: Set(true), + generated_hostname_mode: Set("standard".to_string()), + sync_generated_records: Set(false), + ..Default::default() + } + .insert(db.as_ref()) + .await + .expect("insert equivalent managed zone"); + } + + let dns_provider_service = + Arc::new(temps_dns::services::DnsProviderService::new(db, encryption)); + let state = Arc::new(DomainAppState { + tls_service: Arc::new(TlsService::new(repository.clone(), certificate_provider)), + repository, + domain_service, + dns_provider_service: Some(dns_provider_service), + audit_service: Arc::new(RecordingAudit::default()), + telemetry: Arc::new(NoopTelemetryReporter), + }); + let mut request = Request::builder() + .method(Method::POST) + .uri(format!("/domains/{}/setup-dns", domain.id)) + .header("content-type", "application/json") + .body(Body::from(format!( + r#"{{"dns_provider_id":{}}}"#, + provider.id + ))) + .unwrap(); + request.extensions_mut().insert(auth(vec![ + Permission::DomainsWrite, + Permission::DnsProvidersWrite, + ])); + request.extensions_mut().insert(metadata()); + + let response = temps_domains::configure_routes() + .with_state(state) + .oneshot(request) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::CONFLICT); + let body = response.into_body().collect().await.unwrap().to_bytes(); + let problem: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!(problem["title"], "Ambiguous Managed DNS Zone"); + assert!(problem["detail"].as_str().unwrap().contains("example.com")); + assert!(!problem["detail"] + .as_str() + .unwrap() + .contains("must-not-be-published")); +} diff --git a/crates/temps-migrations/src/lib.rs b/crates/temps-migrations/src/lib.rs index 5bd34fe46..037bc59d5 100644 --- a/crates/temps-migrations/src/lib.rs +++ b/crates/temps-migrations/src/lib.rs @@ -9,4 +9,5 @@ pub use sea_orm_migration::prelude::*; mod migration; // Re-export for convenience // Re-export removed +pub use migration::m20260805_000001_index_normalized_managed_domains::Migration as NormalizedManagedDomainIndexMigration; pub use migration::Migrator; diff --git a/crates/temps-migrations/src/migration/m20260805_000001_index_normalized_managed_domains.rs b/crates/temps-migrations/src/migration/m20260805_000001_index_normalized_managed_domains.rs new file mode 100644 index 000000000..533afd050 --- /dev/null +++ b/crates/temps-migrations/src/migration/m20260805_000001_index_normalized_managed_domains.rs @@ -0,0 +1,55 @@ +//! Indexes managed DNS zones by the canonical form used for authoritative +//! suffix lookup. Existing rows may contain mixed case, surrounding whitespace, +//! a wildcard prefix, or a trailing root dot, so canonicalization stays in the +//! index rather than requiring an unsafe data rewrite. + +use sea_orm_migration::prelude::*; + +const INDEX_NAME: &str = "idx_dns_managed_domains_normalized_domain"; +const NORMALIZED_DOMAIN_SQL: &str = + "LOWER(REGEXP_REPLACE(RTRIM(BTRIM(\"domain\"), '.'), '^((\\*\\.)+)', ''))"; + +#[derive(DeriveMigrationName)] +pub struct Migration; + +#[async_trait::async_trait] +impl MigrationTrait for Migration { + async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> { + // SeaORM 1.1 wraps PostgreSQL migrations in a transaction, so + // CREATE INDEX CONCURRENTLY is not legal here. lock_timeout only + // bounds lock acquisition; statement_timeout also bounds the complete + // index operation. This configuration table is intentionally small, + // so a short non-concurrent build is preferable to an unbounded wait. + manager + .get_connection() + .execute_unprepared("SET LOCAL lock_timeout = '5s'") + .await?; + manager + .get_connection() + .execute_unprepared("SET LOCAL statement_timeout = '30s'") + .await?; + manager + .get_connection() + .execute_unprepared(&format!( + "CREATE INDEX IF NOT EXISTS {INDEX_NAME} ON dns_managed_domains (({NORMALIZED_DOMAIN_SQL}))" + )) + .await?; + Ok(()) + } + + async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> { + manager + .get_connection() + .execute_unprepared("SET LOCAL lock_timeout = '5s'") + .await?; + manager + .get_connection() + .execute_unprepared("SET LOCAL statement_timeout = '30s'") + .await?; + manager + .get_connection() + .execute_unprepared(&format!("DROP INDEX IF EXISTS {INDEX_NAME}")) + .await?; + Ok(()) + } +} diff --git a/crates/temps-migrations/src/migration/mod.rs b/crates/temps-migrations/src/migration/mod.rs index a01d3830d..5d6ad6846 100644 --- a/crates/temps-migrations/src/migration/mod.rs +++ b/crates/temps-migrations/src/migration/mod.rs @@ -173,6 +173,7 @@ mod m20260803_000001_add_flag_last_evaluated_at; mod m20260803_000001_add_template_slug_to_projects; mod m20260803_000002_add_step_up_expires_at_to_sessions; mod m20260804_000001_add_must_change_password_to_users; +pub mod m20260805_000001_index_normalized_managed_domains; pub struct Migrator; @@ -353,6 +354,7 @@ impl MigratorTrait for Migrator { Box::new(m20260803_000001_add_template_slug_to_projects::Migration), Box::new(m20260803_000002_add_step_up_expires_at_to_sessions::Migration), Box::new(m20260804_000001_add_must_change_password_to_users::Migration), + Box::new(m20260805_000001_index_normalized_managed_domains::Migration), ] } } diff --git a/crates/temps-migrations/tests/normalized_managed_domain_index_test.rs b/crates/temps-migrations/tests/normalized_managed_domain_index_test.rs new file mode 100644 index 000000000..1b1fcc046 --- /dev/null +++ b/crates/temps-migrations/tests/normalized_managed_domain_index_test.rs @@ -0,0 +1,198 @@ +use sea_orm::{ConnectionTrait, Database, DatabaseConnection, Statement, TransactionTrait}; +use sea_orm_migration::{MigrationTrait, MigratorTrait, SchemaManager}; +use temps_migrations::Migrator; +use temps_migrations::NormalizedManagedDomainIndexMigration; +use testcontainers::{runners::AsyncRunner, GenericImage, ImageExt}; + +const INDEX_NAME: &str = "idx_dns_managed_domains_normalized_domain"; + +async fn connect_with_retries(database_url: &str) -> anyhow::Result { + let mut retries = 5; + loop { + match Database::connect(database_url).await { + Ok(db) => return Ok(db), + Err(error) if retries > 0 => { + retries -= 1; + tokio::time::sleep(tokio::time::Duration::from_secs(2)).await; + if retries == 0 { + return Err(error.into()); + } + } + Err(error) => return Err(error.into()), + } + } +} + +async fn index_exists(db: &DatabaseConnection) -> anyhow::Result { + let row = db + .query_one(Statement::from_string( + sea_orm::DatabaseBackend::Postgres, + format!("SELECT to_regclass('public.{INDEX_NAME}') IS NOT NULL AS present"), + )) + .await? + .expect("index existence query returns one row"); + Ok(row.try_get("", "present")?) +} + +async fn index_definition(db: &C) -> anyhow::Result +where + C: ConnectionTrait, +{ + let row = db + .query_one(Statement::from_string( + sea_orm::DatabaseBackend::Postgres, + format!("SELECT pg_get_indexdef('{INDEX_NAME}'::regclass) AS definition"), + )) + .await? + .expect("index definition query returns one row"); + Ok(row.try_get("", "definition")?) +} + +#[tokio::test] +async fn test_normalized_managed_domain_index_migration_is_used_and_reversible( +) -> anyhow::Result<()> { + if std::env::var("TEMPS_TEST_DATABASE_URL") + .map(|value| !value.trim().is_empty()) + .unwrap_or(false) + { + eprintln!("Skipping normalized-domain index migration test: external database in use"); + return Ok(()); + } + + let container = match GenericImage::new("timescale/timescaledb-ha", "pg18") + .with_env_var("POSTGRES_DB", "postgres") + .with_env_var("POSTGRES_USER", "postgres") + .with_env_var("POSTGRES_PASSWORD", "postgres") + .with_env_var("POSTGRES_HOST_AUTH_METHOD", "trust") + .with_cmd(vec![ + "postgres", + "-c", + "timescaledb.max_background_workers=0", + ]) + .start() + .await + { + Ok(container) => container, + Err(error) => { + eprintln!( + "Skipping normalized-domain index migration test: Docker unavailable: {error}" + ); + return Ok(()); + } + }; + let port = container.get_host_port_ipv4(5432).await?; + let database_url = format!("postgresql://postgres:postgres@localhost:{port}/postgres"); + tokio::time::sleep(tokio::time::Duration::from_secs(3)).await; + let db = connect_with_retries(&database_url).await?; + + Migrator::up(&db, None).await?; + assert!( + index_exists(&db).await?, + "normalized-domain index must exist after up" + ); + + // Invoke this migration directly even though the migrator already created + // the index. Both calls execute the CREATE INDEX IF NOT EXISTS statement. + let transaction = db.begin().await?; + let manager = SchemaManager::new(&transaction); + NormalizedManagedDomainIndexMigration.up(&manager).await?; + NormalizedManagedDomainIndexMigration.up(&manager).await?; + assert_eq!( + transaction + .query_one(Statement::from_string( + sea_orm::DatabaseBackend::Postgres, + "SELECT current_setting('lock_timeout') AS lock_timeout, \ + current_setting('statement_timeout') AS statement_timeout" + .to_string(), + )) + .await? + .expect("timeout settings query returns one row") + .try_get::("", "lock_timeout")?, + "5s" + ); + let timeout_row = transaction + .query_one(Statement::from_string( + sea_orm::DatabaseBackend::Postgres, + "SELECT current_setting('statement_timeout') AS statement_timeout".to_string(), + )) + .await? + .expect("statement timeout query returns one row"); + assert_eq!( + timeout_row.try_get::("", "statement_timeout")?, + "30s" + ); + transaction.commit().await?; + assert!( + index_exists(&db).await?, + "direct repeated up must preserve the index" + ); + let definition = index_definition(&db).await?; + assert!( + definition.contains( + "lower(regexp_replace(rtrim(btrim((domain)::text), '.'::text), '^((\\*\\.)+)'::text, ''::text))" + ), + "index must retain the canonical managed-domain expression; definition was: {definition}" + ); + + db.execute_unprepared( + "INSERT INTO dns_providers (name, provider_type, credentials) \ + VALUES ('index-plan-provider', 'manual', 'unused')", + ) + .await?; + db.execute_unprepared( + "INSERT INTO dns_managed_domains (provider_id, domain, verified) \ + SELECT (SELECT id FROM dns_providers WHERE name = 'index-plan-provider'), \ + 'zone-' || value || '.example.test', true \ + FROM generate_series(1, 2000) AS value", + ) + .await?; + db.execute_unprepared("SET enable_seqscan = off").await?; + let plan_rows = db + .query_all(Statement::from_string( + sea_orm::DatabaseBackend::Postgres, + "EXPLAIN SELECT * FROM dns_managed_domains \ + WHERE LOWER(REGEXP_REPLACE(RTRIM(BTRIM(\"dns_managed_domains\".\"domain\"), '.'), '^((\\*\\.)+)', '')) \ + IN ('zone-1500.example.test')" + .to_string(), + )) + .await?; + let plan = plan_rows + .iter() + .map(|row| row.try_get::("", "QUERY PLAN")) + .collect::, _>>()? + .join("\n"); + assert!( + plan.contains(INDEX_NAME), + "planner must use {INDEX_NAME}; plan was:\n{plan}" + ); + + let transaction = db.begin().await?; + let manager = SchemaManager::new(&transaction); + NormalizedManagedDomainIndexMigration.down(&manager).await?; + let absent = transaction + .query_one(Statement::from_string( + sea_orm::DatabaseBackend::Postgres, + format!("SELECT to_regclass('public.{INDEX_NAME}') IS NULL AS absent"), + )) + .await? + .expect("index absence query returns one row"); + assert!( + absent.try_get::("", "absent")?, + "direct down must remove the normalized-domain index" + ); + NormalizedManagedDomainIndexMigration.up(&manager).await?; + assert!( + transaction + .query_one(Statement::from_string( + sea_orm::DatabaseBackend::Postgres, + format!("SELECT to_regclass('public.{INDEX_NAME}') IS NOT NULL AS present"), + )) + .await? + .expect("restored index query returns one row") + .try_get::("", "present")?, + "direct up after direct down must restore the normalized-domain index" + ); + transaction.commit().await?; + + Ok(()) +}