From 6e5f330ff56f6e31976ee28f065aaea8711ddb25 Mon Sep 17 00:00:00 2001 From: Zhengchao An Date: Thu, 30 Jul 2026 19:20:17 +0800 Subject: [PATCH] feat(kms): add operation timeout and typed retry policy engine (#5472) --- Cargo.lock | 3 + Cargo.toml | 1 + crates/kms/Cargo.toml | 6 + crates/kms/src/backends/vault.rs | 19 +- crates/kms/src/backends/vault_transit.rs | 20 +- crates/kms/src/config.rs | 71 ++- crates/kms/src/error.rs | 18 + crates/kms/src/lib.rs | 4 + crates/kms/src/policy.rs | 605 +++++++++++++++++++++++ 9 files changed, 735 insertions(+), 12 deletions(-) create mode 100644 crates/kms/src/policy.rs diff --git a/Cargo.lock b/Cargo.lock index a2c01b043..0acca68ae 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9441,6 +9441,7 @@ name = "rustfs-kms" version = "1.0.0-beta.12" dependencies = [ "aes-gcm", + "anyhow", "arc-swap", "argon2", "async-trait", @@ -9455,6 +9456,7 @@ dependencies = [ "reqwest", "rustfs-security-governance", "rustfs-utils", + "rustify", "serde", "serde_json", "sha2 0.11.0", @@ -9463,6 +9465,7 @@ dependencies = [ "tempfile", "thiserror 2.0.19", "tokio", + "tokio-util", "tracing", "url", "uuid", diff --git a/Cargo.toml b/Cargo.toml index d8512078f..34a4b6327 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -285,6 +285,7 @@ reed-solomon-simd = "3.1.0" regex = { version = "1.13.1" } rumqttc = { package = "rumqttc-next", version = "0.33.3" } redis = { version = "1.5.0" } +rustify = { version = "0.7", default-features = false } rustix = { version = "1.1.4" } rust-embed = { version = "8.12.0" } rustc-hash = { version = "2.1.3" } diff --git a/crates/kms/Cargo.toml b/crates/kms/Cargo.toml index 370f1e90d..e492f0e83 100644 --- a/crates/kms/Cargo.toml +++ b/crates/kms/Cargo.toml @@ -65,11 +65,17 @@ rustfs-security-governance = { workspace = true } # HTTP client for Vault reqwest = { workspace = true } vaultrs = { workspace = true } +# vaultrs surfaces transport-level failures as wrapped rustify errors; the +# operation policy needs the concrete type to classify them for retry decisions. +rustify = { workspace = true } +tokio-util = { workspace = true } [dev-dependencies] +anyhow = { workspace = true } insta = { workspace = true, features = ["yaml", "json"] } tempfile = { workspace = true } temp-env = { workspace = true } +tokio = { workspace = true, features = ["test-util"] } [features] default = [] diff --git a/crates/kms/src/backends/vault.rs b/crates/kms/src/backends/vault.rs index 0ba44b23b..b9f8e6364 100644 --- a/crates/kms/src/backends/vault.rs +++ b/crates/kms/src/backends/vault.rs @@ -68,10 +68,17 @@ struct VaultKeyData { impl VaultKmsClient { /// Create a new Vault KMS client - pub async fn new(config: VaultConfig) -> Result { + /// + /// `attempt_timeout` caps every HTTP request issued through this client. + pub async fn new(config: VaultConfig, attempt_timeout: Duration) -> Result { // Create client settings let mut settings_builder = VaultClientSettingsBuilder::default(); settings_builder.address(&config.address); + // Defense in depth against stalled connections: vaultrs leaves the + // underlying reqwest client without any timeout by default, so a hung + // request would otherwise wait forever regardless of the + // operation-level retry policy. + settings_builder.timeout(Some(attempt_timeout)); // Set authentication token based on method let token = match &config.auth_method { @@ -607,7 +614,7 @@ impl VaultKmsBackend { } }; - let client = VaultKmsClient::new(vault_config).await?; + let client = VaultKmsClient::new(vault_config, config.effective_timeout()).await?; Ok(Self { client }) } @@ -851,7 +858,9 @@ mod tests { tls: None, }; - let client = VaultKmsClient::new(config).await.expect("Failed to create Vault client"); + let client = VaultKmsClient::new(config, Duration::from_secs(30)) + .await + .expect("Failed to create Vault client"); // Test key operations let key_id = "test-key-vault"; @@ -906,7 +915,9 @@ mod tests { // Regression: get_key_material previously "self-healed" a decrypt/length failure by // minting a fresh random master key and overwriting the stored value — destroying the // original key and making every DEK wrapped by it permanently undecryptable. - let client = VaultKmsClient::new(integration_vault_config()).await.expect("client"); + let client = VaultKmsClient::new(integration_vault_config(), Duration::from_secs(30)) + .await + .expect("client"); let key_id = format!("corrupt-{}", uuid::Uuid::new_v4()); client.create_key(&key_id, "AES_256", None).await.expect("create"); diff --git a/crates/kms/src/backends/vault_transit.rs b/crates/kms/src/backends/vault_transit.rs index 7be6487ba..9b95f826f 100644 --- a/crates/kms/src/backends/vault_transit.rs +++ b/crates/kms/src/backends/vault_transit.rs @@ -138,9 +138,17 @@ pub struct VaultTransitKmsClient { } impl VaultTransitKmsClient { - pub async fn new(config: VaultTransitConfig) -> Result { + /// Create a new Vault Transit KMS client + /// + /// `attempt_timeout` caps every HTTP request issued through this client. + pub async fn new(config: VaultTransitConfig, attempt_timeout: Duration) -> Result { let mut settings_builder = VaultClientSettingsBuilder::default(); settings_builder.address(&config.address); + // Defense in depth against stalled connections: vaultrs leaves the + // underlying reqwest client without any timeout by default, so a hung + // request would otherwise wait forever regardless of the + // operation-level retry policy. + settings_builder.timeout(Some(attempt_timeout)); let token = match &config.auth_method { crate::config::VaultAuthMethod::Token { token } => token.clone(), @@ -607,7 +615,7 @@ impl VaultTransitKmsBackend { } }; - let client = VaultTransitKmsClient::new(vault_config).await?; + let client = VaultTransitKmsClient::new(vault_config, config.effective_timeout()).await?; Ok(Self { client }) } } @@ -788,7 +796,7 @@ mod tests { let config = test_vault_transit_config(); // --- First "process": create a key and disable it --- - let client1 = VaultTransitKmsClient::new(config.clone()) + let client1 = VaultTransitKmsClient::new(config.clone(), Duration::from_secs(30)) .await .expect("Failed to create VaultTransit client"); @@ -811,7 +819,7 @@ mod tests { assert_eq!(info_after_disable.status, KeyStatus::Disabled, "key must be Disabled after disable_key"); // --- Simulate restart: create a brand new client with empty cache --- - let client2 = VaultTransitKmsClient::new(config) + let client2 = VaultTransitKmsClient::new(config, Duration::from_secs(30)) .await .expect("Failed to create second VaultTransit client (restart simulation)"); @@ -841,7 +849,7 @@ mod tests { async fn test_transit_pending_deletion_survives_restart_simulation() { let config = test_vault_transit_config(); - let client1 = VaultTransitKmsClient::new(config.clone()) + let client1 = VaultTransitKmsClient::new(config.clone(), Duration::from_secs(30)) .await .expect("Failed to create VaultTransit client"); @@ -865,7 +873,7 @@ mod tests { "key must be PendingDeletion after schedule_key_deletion" ); - let client2 = VaultTransitKmsClient::new(config) + let client2 = VaultTransitKmsClient::new(config, Duration::from_secs(30)) .await .expect("Failed to create second VaultTransit client (restart simulation)"); diff --git a/crates/kms/src/config.rs b/crates/kms/src/config.rs index 3d2a469a5..04aab91d6 100644 --- a/crates/kms/src/config.rs +++ b/crates/kms/src/config.rs @@ -32,6 +32,15 @@ pub const ENV_KMS_STATIC_SECRET_KEY_FILE: &str = "RUSTFS_KMS_STATIC_SECRET_KEY_F pub const DEFAULT_VAULT_TRANSIT_METADATA_KV_MOUNT: &str = "secret"; pub const DEFAULT_VAULT_TRANSIT_METADATA_KEY_PREFIX: &str = "rustfs/kms/transit-metadata"; +/// Upper bound applied to `KmsConfig::timeout` when deriving backend behavior. +/// +/// Out-of-range values are clamped at use rather than rejected so existing +/// deployments with oversized settings keep starting after an upgrade. +pub(crate) const MAX_OPERATION_TIMEOUT: Duration = Duration::from_secs(300); + +/// Upper bound applied to `KmsConfig::retry_attempts` when deriving backend behavior. +pub(crate) const MAX_RETRY_ATTEMPTS: u32 = 10; + fn default_vault_transit_metadata_kv_mount() -> String { DEFAULT_VAULT_TRANSIT_METADATA_KV_MOUNT.to_string() } @@ -117,9 +126,15 @@ pub struct KmsConfig { /// Allow development-only insecure defaults such as plaintext local keys or HTTP Vault. #[serde(default)] pub allow_insecure_dev_defaults: bool, - /// Operation timeout + /// Timeout for a single backend attempt. + /// + /// This bounds one outbound request, not the whole operation: the operation + /// policy owns the total deadline across retries. Values above 300 seconds + /// are clamped at use (see `KmsConfig::effective_timeout`). pub timeout: Duration, - /// Number of retry attempts + /// Number of retry attempts. + /// + /// Values above 10 are clamped at use (see `KmsConfig::effective_retry_attempts`). pub retry_attempts: u32, /// Enable caching pub enable_cache: bool, @@ -536,6 +551,16 @@ impl KmsConfig { self } + /// Per-attempt timeout with the configured value clamped to the supported maximum. + pub(crate) fn effective_timeout(&self) -> Duration { + self.timeout.min(MAX_OPERATION_TIMEOUT) + } + + /// Retry attempts with the configured value clamped to the supported maximum. + pub(crate) fn effective_retry_attempts(&self) -> u32 { + self.retry_attempts.min(MAX_RETRY_ATTEMPTS) + } + /// Validate the configuration pub fn validate(&self) -> Result<()> { // Validate timeout @@ -548,6 +573,23 @@ impl KmsConfig { return Err(KmsError::configuration_error("Retry attempts must be greater than 0")); } + // Oversized values are clamped at use (not rejected) so pre-existing + // configurations cannot keep the service from starting after upgrade. + if self.timeout > MAX_OPERATION_TIMEOUT { + tracing::warn!( + configured_secs = self.timeout.as_secs(), + max_secs = MAX_OPERATION_TIMEOUT.as_secs(), + "KMS timeout exceeds the supported maximum; backend operations clamp it to the maximum" + ); + } + if self.retry_attempts > MAX_RETRY_ATTEMPTS { + tracing::warn!( + configured = self.retry_attempts, + max = MAX_RETRY_ATTEMPTS, + "KMS retry_attempts exceeds the supported maximum; backend operations clamp it to the maximum" + ); + } + // Validate backend-specific configuration match &self.backend_config { BackendConfig::Local(config) => { @@ -855,6 +897,31 @@ mod tests { assert_eq!(local_config.key_dir, temp_dir.path()); } + #[test] + fn test_oversized_timeout_and_retries_clamped_not_rejected() { + let temp_dir = TempDir::new().expect("Failed to create temp dir"); + let config = KmsConfig { + timeout: Duration::from_secs(3_600), + retry_attempts: 50, + ..KmsConfig::local(temp_dir.path().to_path_buf()).with_insecure_development_defaults() + }; + + // Out-of-range values must not keep the service from starting. + assert!(config.validate().is_ok()); + assert_eq!(config.effective_timeout(), MAX_OPERATION_TIMEOUT); + assert_eq!(config.effective_retry_attempts(), MAX_RETRY_ATTEMPTS); + + // In-range values pass through unchanged. + let config = KmsConfig { + timeout: Duration::from_secs(45), + retry_attempts: 5, + ..config + }; + assert!(config.validate().is_ok()); + assert_eq!(config.effective_timeout(), Duration::from_secs(45)); + assert_eq!(config.effective_retry_attempts(), 5); + } + #[test] fn test_local_development_defaults_require_opt_in() { let temp_dir = TempDir::new().expect("Failed to create temp dir"); diff --git a/crates/kms/src/error.rs b/crates/kms/src/error.rs index f6999aa51..d1031a32f 100644 --- a/crates/kms/src/error.rs +++ b/crates/kms/src/error.rs @@ -90,6 +90,14 @@ pub enum KmsError { /// Encryption context mismatch #[error("Encryption context mismatch: {message}")] ContextMismatch { message: String }, + + /// Backend operation exceeded its per-attempt timeout or total deadline + #[error("Operation timed out: {message}")] + OperationTimedOut { message: String }, + + /// Backend operation aborted by cancellation or shutdown + #[error("Operation cancelled: {message}")] + OperationCancelled { message: String }, } impl KmsError { @@ -187,6 +195,16 @@ impl KmsError { pub fn context_mismatch>(message: S) -> Self { Self::ContextMismatch { message: message.into() } } + + /// Create an operation timed out error + pub fn operation_timed_out>(message: S) -> Self { + Self::OperationTimedOut { message: message.into() } + } + + /// Create an operation cancelled error + pub fn operation_cancelled>(message: S) -> Self { + Self::OperationCancelled { message: message.into() } + } } /// Convert from standard library errors diff --git a/crates/kms/src/lib.rs b/crates/kms/src/lib.rs index db4791ec3..d15718ef9 100644 --- a/crates/kms/src/lib.rs +++ b/crates/kms/src/lib.rs @@ -71,6 +71,10 @@ pub mod config; mod encryption; mod error; pub mod manager; +// The executor is wired into the Vault backends in a follow-up change; until +// then the module is only exercised by its own tests. +#[allow(dead_code)] +mod policy; pub mod service; pub mod service_manager; mod time_serde; diff --git a/crates/kms/src/policy.rs b/crates/kms/src/policy.rs new file mode 100644 index 000000000..c2014b0e0 --- /dev/null +++ b/crates/kms/src/policy.rs @@ -0,0 +1,605 @@ +// Copyright 2024 RustFS Team +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Operation execution policy for external KMS backend calls. +//! +//! Vault-backed operations leave the process boundary, so every call needs a +//! per-attempt timeout, a total operation deadline, and classification-driven +//! bounded retries. This module provides the engine only; the Vault backends +//! wire their call sites through [`execute`] in a follow-up change. +//! +//! Retry safety is driven by two orthogonal classifications: +//! - [`OpClass`] states whether replaying the operation is safe at all. +//! - [`ErrorClass`] states whether the observed failure is worth replaying. +//! +//! Mutating operations without an idempotency key or CAS precondition are never +//! retried automatically: a response lost after the server applied the write +//! would otherwise be replayed into duplicate side effects (extra key versions, +//! repeated deletes). + +use std::future::Future; +use std::time::Duration; + +use rand::{RngExt, SeedableRng, rngs::StdRng}; +use tokio::time::Instant; +use tokio_util::sync::CancellationToken; + +use crate::config::KmsConfig; +use crate::error::{KmsError, Result}; + +/// Default backoff cap before the first retry; doubles per retry. +const DEFAULT_BASE_BACKOFF: Duration = Duration::from_millis(100); + +/// Default upper bound for a single backoff sleep. +const DEFAULT_MAX_BACKOFF: Duration = Duration::from_secs(2); + +/// Replay safety of a backend operation. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum OpClass { + /// No persistent side effects (encrypt/decrypt/generate/describe/list/health); + /// safe to retry on any retryable failure. + ReadIdempotent, + /// External write without an idempotency key or CAS precondition + /// (create/rotate/delete/configure); executed at most once. + MutatingNonIdempotent, + /// Authentication exchange (login, token renewal); safe to resend. + Auth, +} + +impl OpClass { + fn retryable(self) -> bool { + !matches!(self, OpClass::MutatingNonIdempotent) + } +} + +/// Retry-relevant classification of a failed attempt. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ErrorClass { + /// Connection-level failure (connect/send/response read); the server may or + /// may not have observed the request. + RetryableConn, + /// Retryable HTTP status: 429 throttling or a recoverable 5xx. + RetryableStatus, + /// Deterministic failure (auth, validation, not-found, malformed data); + /// retrying cannot help and may mask the real problem. + Fatal, +} + +/// Classify a `vaultrs` client error for retry purposes. +/// +/// Status codes are inspected both on `ClientError::APIError` (JSON error body) +/// and on a wrapped rustify `ServerResponseError` (non-JSON body, e.g. an HTML +/// page from an intermediate load balancer). Everything that is not throttling, +/// a recoverable 5xx, or a connection-level failure is fatal; in particular +/// 400/401/403/404 must never be retried. +pub(crate) fn classify_vaultrs(error: &vaultrs::error::ClientError) -> ErrorClass { + use rustify::errors::ClientError as RestError; + use vaultrs::error::ClientError; + + match error { + ClientError::APIError { code, .. } => classify_status(*code), + ClientError::RestClientError { source } => match source { + RestError::ServerResponseError { code, .. } => classify_status(*code), + RestError::RequestError { .. } | RestError::ResponseError { .. } => ErrorClass::RetryableConn, + _ => ErrorClass::Fatal, + }, + _ => ErrorClass::Fatal, + } +} + +fn classify_status(code: u16) -> ErrorClass { + match code { + 429 | 500 | 502 | 503 | 504 => ErrorClass::RetryableStatus, + _ => ErrorClass::Fatal, + } +} + +/// Failure of a single attempt, carrying its retry classification. +#[derive(Debug)] +pub(crate) struct AttemptError { + pub(crate) class: ErrorClass, + pub(crate) error: KmsError, +} + +/// Budgets applied by [`execute`]. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct RetryPolicy { + /// Upper bound for one backend attempt. + pub(crate) attempt_timeout: Duration, + /// Upper bound for the whole operation, including retries and backoff. + pub(crate) op_deadline: Duration, + /// Maximum attempts for retryable operation classes; non-idempotent + /// mutations always run exactly once regardless of this value. + pub(crate) max_attempts: u32, + /// Backoff cap before the first retry; doubles per retry. + pub(crate) base_backoff: Duration, + /// Upper bound for a single backoff sleep. + pub(crate) max_backoff: Duration, +} + +impl RetryPolicy { + /// Derive the policy from the KMS configuration. + /// + /// `timeout` and `retry_attempts` are taken with the config-level clamps + /// applied. The operation deadline covers the worst-case budget of all + /// attempts plus backoff, so it bounds runaway loops without cutting any + /// attempt short; an independently configurable deadline is left to the + /// admin-API follow-up. + pub(crate) fn from_config(config: &KmsConfig) -> Self { + let mut policy = Self { + attempt_timeout: config.effective_timeout(), + op_deadline: Duration::ZERO, + max_attempts: config.effective_retry_attempts(), + base_backoff: DEFAULT_BASE_BACKOFF, + max_backoff: DEFAULT_MAX_BACKOFF, + }; + policy.op_deadline = policy.worst_case_budget(); + policy + } + + /// Total worst-case duration: every attempt hits `attempt_timeout` and + /// every backoff sleeps its full cap. + fn worst_case_budget(&self) -> Duration { + let attempts = self.max_attempts.max(1); + let mut budget = self.attempt_timeout.saturating_mul(attempts); + for completed in 1..attempts { + budget = budget.saturating_add(backoff_cap(self, completed)); + } + budget + } +} + +/// Exponential backoff cap after `completed_attempts` failed attempts. +fn backoff_cap(policy: &RetryPolicy, completed_attempts: u32) -> Duration { + let doublings = completed_attempts.saturating_sub(1).min(31); + policy.base_backoff.saturating_mul(1u32 << doublings).min(policy.max_backoff) +} + +/// Equal jitter: sleep within `[cap / 2, cap]`, keeping at least half the cap +/// so backoff still backs off while decorrelating retry bursts. +fn equal_jitter(rng: &mut impl RngExt, cap: Duration) -> Duration { + let half = cap / 2; + let spread = u64::try_from(half.as_nanos()).unwrap_or(u64::MAX); + half + Duration::from_nanos(rng.random_range(0..=spread)) +} + +/// Run `attempt` under the policy. +/// +/// Each attempt is bounded by `attempt_timeout` (further capped by whatever is +/// left of `op_deadline`), and retryable failures are replayed with exponential +/// backoff and jitter when the operation class allows it. Cancellation aborts +/// both in-flight attempts and backoff sleeps. +/// +/// A timed-out attempt counts as a connection-class failure: the server may +/// have processed the request, which is exactly why non-idempotent mutations +/// are never replayed. +pub(crate) async fn execute( + operation: &'static str, + class: OpClass, + policy: &RetryPolicy, + cancel: &CancellationToken, + attempt: F, +) -> Result +where + F: FnMut() -> Fut, + Fut: Future>, +{ + // Seed an owned RNG up front: the thread-local RNG is not Send and must + // not be held across await points. + let mut rng = StdRng::from_rng(&mut rand::rng()); + execute_with_jitter(operation, class, policy, cancel, move |cap| equal_jitter(&mut rng, cap), attempt).await +} + +/// [`execute`] with an injectable jitter source so tests can pin deterministic +/// backoff durations instead of asserting around random sleeps. +pub(crate) async fn execute_with_jitter( + operation: &'static str, + class: OpClass, + policy: &RetryPolicy, + cancel: &CancellationToken, + mut jitter: J, + mut attempt: F, +) -> Result +where + F: FnMut() -> Fut, + Fut: Future>, + J: FnMut(Duration) -> Duration, +{ + let deadline = Instant::now() + policy.op_deadline; + let max_attempts = if class.retryable() { policy.max_attempts.max(1) } else { 1 }; + + let mut attempt_no = 0u32; + loop { + attempt_no += 1; + if cancel.is_cancelled() { + return Err(KmsError::operation_cancelled(format!( + "{operation} cancelled before attempt {attempt_no}" + ))); + } + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + return Err(KmsError::operation_timed_out(format!( + "{operation} exceeded operation deadline of {:?}", + policy.op_deadline + ))); + } + + let attempt_budget = policy.attempt_timeout.min(remaining); + let outcome = tokio::select! { + biased; + _ = cancel.cancelled() => { + return Err(KmsError::operation_cancelled(format!("{operation} cancelled during attempt {attempt_no}"))); + } + outcome = tokio::time::timeout(attempt_budget, attempt()) => outcome, + }; + + let failure = match outcome { + Ok(Ok(value)) => return Ok(value), + Ok(Err(failure)) => failure, + Err(_) => AttemptError { + class: ErrorClass::RetryableConn, + error: KmsError::operation_timed_out(format!( + "{operation} attempt {attempt_no} timed out after {attempt_budget:?}" + )), + }, + }; + + if failure.class == ErrorClass::Fatal || attempt_no >= max_attempts { + return Err(failure.error); + } + + let backoff = jitter(backoff_cap(policy, attempt_no)); + if backoff >= deadline.saturating_duration_since(Instant::now()) { + // Not enough deadline budget left for another attempt. + return Err(failure.error); + } + tracing::warn!( + operation, + attempt = attempt_no, + error_class = ?failure.class, + backoff = ?backoff, + "KMS backend attempt failed with a retryable error; backing off before retry" + ); + tokio::select! { + biased; + _ = cancel.cancelled() => { + return Err(KmsError::operation_cancelled(format!("{operation} cancelled during retry backoff"))); + } + _ = tokio::time::sleep(backoff) => {} + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Arc; + use std::sync::atomic::{AtomicU32, Ordering}; + + type AttemptResult = std::result::Result; + + fn policy_of(attempt_timeout_ms: u64, op_deadline_ms: u64, max_attempts: u32, base_ms: u64, max_ms: u64) -> RetryPolicy { + RetryPolicy { + attempt_timeout: Duration::from_millis(attempt_timeout_ms), + op_deadline: Duration::from_millis(op_deadline_ms), + max_attempts, + base_backoff: Duration::from_millis(base_ms), + max_backoff: Duration::from_millis(max_ms), + } + } + + /// Deterministic jitter: always sleep the full backoff cap. + fn full_jitter(cap: Duration) -> Duration { + cap + } + + fn retryable_conn_error() -> AttemptError { + AttemptError { + class: ErrorClass::RetryableConn, + error: KmsError::backend_error("connection reset by peer"), + } + } + + #[tokio::test(start_paused = true)] + async fn hung_attempt_fails_within_attempt_timeout() { + let policy = policy_of(5_000, 60_000, 1, 100, 2_000); + let cancel = CancellationToken::new(); + let started = Instant::now(); + + let result: Result<()> = execute_with_jitter("encrypt", OpClass::ReadIdempotent, &policy, &cancel, full_jitter, || { + std::future::pending::>() + }) + .await; + + assert!(matches!(result, Err(KmsError::OperationTimedOut { .. })), "got {result:?}"); + assert_eq!(started.elapsed(), Duration::from_millis(5_000)); + } + + #[tokio::test(start_paused = true)] + async fn retryable_status_retries_until_success() { + let policy = policy_of(1_000, 60_000, 3, 100, 2_000); + let cancel = CancellationToken::new(); + let calls = Arc::new(AtomicU32::new(0)); + let calls_in_attempt = calls.clone(); + let started = Instant::now(); + + let result = + execute_with_jitter("generate_data_key", OpClass::ReadIdempotent, &policy, &cancel, full_jitter, move || { + let calls = calls_in_attempt.clone(); + async move { + if calls.fetch_add(1, Ordering::SeqCst) < 2 { + Err(AttemptError { + class: ErrorClass::RetryableStatus, + error: KmsError::backend_error("throttled (429)"), + }) + } else { + Ok(7u32) + } + } + }) + .await; + + assert_eq!(result.expect("retries within budget must succeed"), 7); + assert_eq!(calls.load(Ordering::SeqCst), 3); + // Full-cap backoff: 100ms after attempt 1, 200ms after attempt 2. + assert_eq!(started.elapsed(), Duration::from_millis(300)); + } + + #[tokio::test(start_paused = true)] + async fn fatal_error_is_not_retried() { + let policy = policy_of(1_000, 60_000, 5, 100, 2_000); + let cancel = CancellationToken::new(); + let calls = Arc::new(AtomicU32::new(0)); + let calls_in_attempt = calls.clone(); + let started = Instant::now(); + + let result: Result<()> = + execute_with_jitter("decrypt", OpClass::ReadIdempotent, &policy, &cancel, full_jitter, move || { + let calls = calls_in_attempt.clone(); + async move { + calls.fetch_add(1, Ordering::SeqCst); + Err(AttemptError { + class: ErrorClass::Fatal, + error: KmsError::access_denied("permission denied (403)"), + }) + } + }) + .await; + + assert!(matches!(result, Err(KmsError::AccessDenied { .. })), "got {result:?}"); + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!(started.elapsed(), Duration::ZERO); + } + + #[tokio::test(start_paused = true)] + async fn mutating_non_idempotent_runs_exactly_once() { + let policy = policy_of(1_000, 60_000, 5, 100, 2_000); + let cancel = CancellationToken::new(); + let calls = Arc::new(AtomicU32::new(0)); + let calls_in_attempt = calls.clone(); + + let result: Result<()> = + execute_with_jitter("rotate_key", OpClass::MutatingNonIdempotent, &policy, &cancel, full_jitter, move || { + let calls = calls_in_attempt.clone(); + async move { + calls.fetch_add(1, Ordering::SeqCst); + Err(retryable_conn_error()) + } + }) + .await; + + // Even a retryable failure must not replay a non-idempotent mutation. + assert!(matches!(result, Err(KmsError::BackendError { .. })), "got {result:?}"); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test(start_paused = true)] + async fn cancel_aborts_backoff_immediately() { + // Long backoff (10s cap) so a prompt return can only come from cancellation. + let policy = policy_of(1_000, 600_000, 5, 10_000, 10_000); + let cancel = CancellationToken::new(); + let canceller = cancel.clone(); + tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(500)).await; + canceller.cancel(); + }); + let calls = Arc::new(AtomicU32::new(0)); + let calls_in_attempt = calls.clone(); + let started = Instant::now(); + + let result: Result<()> = + execute_with_jitter("decrypt", OpClass::ReadIdempotent, &policy, &cancel, full_jitter, move || { + let calls = calls_in_attempt.clone(); + async move { + calls.fetch_add(1, Ordering::SeqCst); + Err(retryable_conn_error()) + } + }) + .await; + + assert!(matches!(result, Err(KmsError::OperationCancelled { .. })), "got {result:?}"); + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!(started.elapsed(), Duration::from_millis(500)); + } + + #[tokio::test(start_paused = true)] + async fn cancelled_token_short_circuits_before_first_attempt() { + let policy = policy_of(1_000, 60_000, 5, 100, 2_000); + let cancel = CancellationToken::new(); + cancel.cancel(); + let calls = Arc::new(AtomicU32::new(0)); + let calls_in_attempt = calls.clone(); + + let result: Result<()> = + execute_with_jitter("list_keys", OpClass::ReadIdempotent, &policy, &cancel, full_jitter, move || { + let calls = calls_in_attempt.clone(); + async move { + calls.fetch_add(1, Ordering::SeqCst); + Ok(()) + } + }) + .await; + + assert!(matches!(result, Err(KmsError::OperationCancelled { .. })), "got {result:?}"); + assert_eq!(calls.load(Ordering::SeqCst), 0); + } + + #[tokio::test(start_paused = true)] + async fn total_duration_never_exceeds_deadline() { + // Worst case without a deadline would be 5 * 10s + backoff; the 25s + // deadline must cut both the attempt count and the final attempt short. + let policy = policy_of(10_000, 25_000, 5, 100, 2_000); + let cancel = CancellationToken::new(); + let calls = Arc::new(AtomicU32::new(0)); + let calls_in_attempt = calls.clone(); + let started = Instant::now(); + + let result: Result<()> = + execute_with_jitter("describe_key", OpClass::ReadIdempotent, &policy, &cancel, full_jitter, move || { + calls_in_attempt.fetch_add(1, Ordering::SeqCst); + std::future::pending::>() + }) + .await; + + assert!(matches!(result, Err(KmsError::OperationTimedOut { .. })), "got {result:?}"); + // attempt 1: 10s + 100ms backoff; attempt 2: 10s + 200ms backoff; + // attempt 3 runs with the residual 4.7s budget, then the deadline is + // exhausted and no further backoff or attempt happens. + assert_eq!(calls.load(Ordering::SeqCst), 3); + assert_eq!(started.elapsed(), Duration::from_millis(25_000)); + } + + #[tokio::test(start_paused = true)] + async fn default_jitter_stays_within_equal_jitter_bounds() { + let policy = policy_of(1_000, 60_000, 3, 100, 2_000); + let cancel = CancellationToken::new(); + let calls = Arc::new(AtomicU32::new(0)); + let calls_in_attempt = calls.clone(); + let started = Instant::now(); + + let result = execute("health_check", OpClass::ReadIdempotent, &policy, &cancel, move || { + let calls = calls_in_attempt.clone(); + async move { + if calls.fetch_add(1, Ordering::SeqCst) < 2 { + Err(retryable_conn_error()) + } else { + Ok(()) + } + } + }) + .await; + + result.expect("retries within budget must succeed"); + let elapsed = started.elapsed(); + // Equal jitter sleeps within [cap / 2, cap]; caps are 100ms then 200ms. + assert!( + elapsed >= Duration::from_millis(150) && elapsed <= Duration::from_millis(300), + "elapsed {elapsed:?} outside equal-jitter bounds" + ); + } + + #[test] + fn equal_jitter_is_seed_deterministic_and_bounded() { + let caps = [Duration::from_millis(100), Duration::from_millis(250), Duration::from_secs(2)]; + let mut first = StdRng::seed_from_u64(1569); + let mut second = StdRng::seed_from_u64(1569); + for cap in caps { + let a = equal_jitter(&mut first, cap); + let b = equal_jitter(&mut second, cap); + assert_eq!(a, b, "same seed must yield the same jitter sequence"); + assert!(a >= cap / 2 && a <= cap, "jitter {a:?} outside [{:?}, {cap:?}]", cap / 2); + } + } + + #[test] + fn backoff_caps_double_up_to_max() { + let policy = policy_of(1_000, 60_000, 6, 100, 700); + let caps: Vec = (1..=5).map(|completed| backoff_cap(&policy, completed)).collect(); + assert_eq!( + caps, + vec![ + Duration::from_millis(100), + Duration::from_millis(200), + Duration::from_millis(400), + Duration::from_millis(700), + Duration::from_millis(700), + ] + ); + } + + #[test] + fn retry_policy_from_config_applies_clamps() { + let config = KmsConfig { + timeout: Duration::from_secs(3_600), + retry_attempts: 50, + ..KmsConfig::default() + }; + let policy = RetryPolicy::from_config(&config); + assert_eq!(policy.attempt_timeout, Duration::from_secs(300)); + assert_eq!(policy.max_attempts, 10); + // The deadline must cover the full worst case so it never cuts a + // policy-conformant operation short. + assert!(policy.op_deadline >= policy.attempt_timeout.saturating_mul(policy.max_attempts)); + + let in_range = KmsConfig::default(); + let policy = RetryPolicy::from_config(&in_range); + assert_eq!(policy.attempt_timeout, in_range.timeout); + assert_eq!(policy.max_attempts, in_range.retry_attempts); + } + + #[test] + fn classify_vaultrs_matrix() { + use rustify::errors::ClientError as RestError; + use vaultrs::error::ClientError; + + let api = |code: u16| ClientError::APIError { code, errors: vec![] }; + for code in [429u16, 500, 502, 503, 504] { + assert_eq!(classify_vaultrs(&api(code)), ErrorClass::RetryableStatus, "status {code}"); + } + for code in [400u16, 401, 403, 404, 405, 412, 501] { + assert_eq!(classify_vaultrs(&api(code)), ErrorClass::Fatal, "status {code}"); + } + + // Status errors whose body could not be parsed stay wrapped in the + // rustify error; classification must still see the code. + let raw_status = |code: u16| ClientError::RestClientError { + source: RestError::ServerResponseError { code, content: None }, + }; + assert_eq!(classify_vaultrs(&raw_status(503)), ErrorClass::RetryableStatus); + assert_eq!(classify_vaultrs(&raw_status(403)), ErrorClass::Fatal); + + let send_failure = ClientError::RestClientError { + source: RestError::RequestError { + source: anyhow::anyhow!("connection refused"), + url: "http://127.0.0.1:8200/v1/sys/health".to_string(), + method: "GET".to_string(), + }, + }; + assert_eq!(classify_vaultrs(&send_failure), ErrorClass::RetryableConn); + + let read_failure = ClientError::RestClientError { + source: RestError::ResponseError { + source: anyhow::anyhow!("connection reset by peer"), + }, + }; + assert_eq!(classify_vaultrs(&read_failure), ErrorClass::RetryableConn); + + let malformed = ClientError::JsonParseError { + source: serde_json::from_str::("{").expect_err("truncated JSON must not parse"), + }; + assert_eq!(classify_vaultrs(&malformed), ErrorClass::Fatal); + assert_eq!(classify_vaultrs(&ClientError::ResponseEmptyError), ErrorClass::Fatal); + assert_eq!(classify_vaultrs(&ClientError::InvalidLoginMethodError), ErrorClass::Fatal); + } +}