diff --git a/TESTING.md b/TESTING.md index a1c2f9476e..c8f5d7630e 100644 --- a/TESTING.md +++ b/TESTING.md @@ -48,6 +48,25 @@ mise run test:rust # cargo test --workspace Rust validation checks tracked Cargo lockfiles; run `mise run rust:lockfiles:check` to check them directly. If one is stale, refresh it with Cargo using its adjacent manifest, review the diff, and commit the update. +### PostgreSQL-backed tests + +Tests that need a real PostgreSQL server, such as advisory-lock concurrency +across two stores, are `#[ignore]`d and named `postgres_*`. Run them with: + +```shell +mise run test:rust:postgres +``` + +The task starts a disposable PostgreSQL container with Docker or Podman +(set `CONTAINER_ENGINE` to choose), runs the tests one at a time, and removes +the container. Each test works in its own temporary schema. To use your own +disposable database, set `OPENSHELL_TEST_POSTGRES_URL`. The task overrides +`OPENSHELL_REPLAY_TEST_DATABASE_URL` so legacy tests use that same database. +Never point it at a database that a running gateway uses: the tests take +fleet-wide advisory locks. +CI does not run these tests; the Kubernetes HA e2e suite covers PostgreSQL end +to end. + ### Native Windows validation Use `mise run --skip-tools pre-commit` with the existing Rust/MSVC toolchain. diff --git a/crates/openshell-server/src/compute/mod.rs b/crates/openshell-server/src/compute/mod.rs index e69fe6e0e2..7e81c9b4d4 100644 --- a/crates/openshell-server/src/compute/mod.rs +++ b/crates/openshell-server/src/compute/mod.rs @@ -5,10 +5,14 @@ pub mod driver_config; pub mod lease; +mod mutation_guard; pub mod provisioning_deadline; mod provisioning_operation; pub mod rootfs_tar; +pub use mutation_guard::MutationScope; +use mutation_guard::{LocalMutationGuard, LocalMutationLocks, MutationGuard}; + use crate::grpc::policy::SANDBOX_SETTINGS_OBJECT_TYPE; use crate::otel_tracing::TraceContextInterceptor; use crate::persistence::{ @@ -217,10 +221,21 @@ impl LifecycleGateRegistry { async fn lock_for(&self, sandbox_id: &str) -> SandboxLifecycleGuard { let gate = self.gate_for(sandbox_id); SandboxLifecycleGuard { + sandbox_id: sandbox_id.to_string(), _guard: gate.lock_owned().await, } } + /// Take the gate only when nobody holds it. Never waits, so a caller may + /// use it while holding the sandbox's local mutation lock. + fn try_lock_for(&self, sandbox_id: &str) -> Option { + let guard = self.gate_for(sandbox_id).try_lock_owned().ok()?; + Some(SandboxLifecycleGuard { + sandbox_id: sandbox_id.to_string(), + _guard: guard, + }) + } + fn gate_for(&self, sandbox_id: &str) -> Arc> { let mut gates = self .gates @@ -248,11 +263,12 @@ impl LifecycleGateRegistry { /// Proof that the current operation holds its sandbox-ID lifecycle gate. /// -/// Lifecycle code must acquire this guard before taking `ComputeRuntime::sync_lock`. -/// Passing it to `lock_global_for_lifecycle` makes that ordering visible at -/// every global-lock acquisition in a lifecycle path. +/// Lifecycle code must acquire this guard before taking the sandbox's local +/// mutation lock (`lock_sandbox_for_lifecycle`). Passing it there makes that +/// ordering visible at every mutation-lock acquisition in a lifecycle path. #[derive(Debug)] pub struct SandboxLifecycleGuard { + sandbox_id: String, _guard: tokio::sync::OwnedMutexGuard<()>, } @@ -669,7 +685,7 @@ pub struct ComputeRuntime { sandbox_watch_bus: SandboxWatchBus, tracing_log_bus: TracingLogBus, supervisor_sessions: Arc, - sync_lock: Arc>, + mutation_locks: Arc, lifecycle_gates: Arc, replica_id: String, /// Gateway-issued staging slots for rootfs tar archives. Shared across @@ -682,13 +698,6 @@ pub struct ComputeRuntime { ssh_identities: Arc>, } -pub struct SandboxSyncGuard { - // Drop the database guard before the local mutex so another local waiter - // cannot race ahead while this replica still owns the cluster-wide lock. - _distributed: crate::persistence::DistributedMutationGuard, - _local: tokio::sync::OwnedMutexGuard<()>, -} - impl fmt::Debug for ComputeRuntime { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.debug_struct("ComputeRuntime").finish_non_exhaustive() @@ -772,7 +781,7 @@ impl ComputeRuntime { sandbox_watch_bus, tracing_log_bus, supervisor_sessions, - sync_lock: Arc::new(Mutex::new(())), + mutation_locks: Arc::new(LocalMutationLocks::new()), lifecycle_gates: Arc::new(LifecycleGateRegistry::default()), replica_id: lease::replica_id(), rootfs_tar_staging, @@ -845,41 +854,6 @@ impl ComputeRuntime { .await } - /// Serializes sandbox/provider-profile invariant checks and object writes - /// across gateway replicas. - /// - /// The local mutex preserves lock ordering within one process. `PostgreSQL` - /// deployments also hold a session-level advisory lock for the duration. - pub(crate) async fn sandbox_sync_guard( - &self, - ) -> crate::persistence::PersistenceResult { - let local = self.sync_lock.clone().lock_owned().await; - let distributed = self.store.acquire_distributed_mutation_guard().await?; - Ok(SandboxSyncGuard { - _distributed: distributed, - _local: local, - }) - } - - pub(crate) async fn sandbox_create_guards( - &self, - sandbox_id: &str, - ) -> crate::persistence::PersistenceResult<(SandboxLifecycleGuard, SandboxSyncGuard)> { - let lifecycle_guard = self.lifecycle_gates.lock_for(sandbox_id).await; - let global_guard = self.sandbox_sync_guard().await?; - Ok((lifecycle_guard, global_guard)) - } - - /// Acquires the process-wide lock for code that already holds the - /// sandbox-ID lifecycle gate. The guard parameter documents and enforces - /// that callers acquire locks in lifecycle-gate -> global-lock order. - async fn lock_global_for_lifecycle( - &self, - _lifecycle_guard: &SandboxLifecycleGuard, - ) -> tokio::sync::OwnedMutexGuard<()> { - self.sync_lock.clone().lock_owned().await - } - #[cfg(test)] pub(crate) fn lifecycle_gate_entry_count(&self) -> usize { self.lifecycle_gates.entry_count() @@ -1127,8 +1101,8 @@ impl ComputeRuntime { launch_authentication: Option>, await_main_process_attachment: bool, ) -> Result { - let (lifecycle_guard, global_guard) = self - .sandbox_create_guards(sandbox.object_id()) + let (lifecycle_guard, mutation_guard) = self + .sandbox_create_guards(sandbox.object_workspace(), sandbox.object_id()) .await .map_err(|error| { crate::grpc::persistence_error_to_status(error, "acquire sandbox mutation lock") @@ -1139,7 +1113,7 @@ impl ComputeRuntime { launch_authentication, await_main_process_attachment, lifecycle_guard, - global_guard, + mutation_guard, )) .await } @@ -1152,7 +1126,7 @@ impl ComputeRuntime { launch_authentication: Option>, await_main_process_attachment: bool, lifecycle_guard: SandboxLifecycleGuard, - global_guard: SandboxSyncGuard, + mutation_guard: MutationGuard, ) -> Result { // Defend the internal create path too, before consuming a staged archive // or persisting the sandbox. The gRPC handler checks before driver validation. @@ -1246,7 +1220,7 @@ impl ComputeRuntime { } } .await; - drop(global_guard); + drop(mutation_guard); let launch_authentication = match prepared { Ok(authentication) => authentication, Err(status) => { @@ -1309,7 +1283,7 @@ impl ComputeRuntime { .await); } }; - let global_guard = self.lock_global_for_lifecycle(&lifecycle_guard).await; + let sandbox_guard = self.lock_sandbox_for_lifecycle(&lifecycle_guard).await; // The scanner can expire preparation while create owns the lifecycle // gate. Every driver outcome must observe that durable decision before // deleting records, publishing status, or compensating a failed create. @@ -1330,7 +1304,7 @@ impl ComputeRuntime { if self.supports_sandbox_authentication() && runtime_identity.is_empty() { let status = Status::internal("compute driver did not return a runtime identity"); return Err(self - .compensate_failed_create(&sandbox, lifecycle_guard, Some(global_guard), status) + .compensate_failed_create(&sandbox, lifecycle_guard, Some(sandbox_guard), status) .await); } if self.supports_sandbox_authentication() { @@ -1359,7 +1333,7 @@ impl ComputeRuntime { .compensate_failed_create( &sandbox, lifecycle_guard, - Some(global_guard), + Some(sandbox_guard), status, ) .await); @@ -1375,7 +1349,7 @@ impl ComputeRuntime { &self, created: &Sandbox, lifecycle_guard: SandboxLifecycleGuard, - global_guard: Option>, + sandbox_guard: Option, original: Status, ) -> Status { // Once create has committed its parent, cancellation must not interrupt @@ -1386,15 +1360,15 @@ impl ComputeRuntime { let request_span = tracing::Span::current(); tokio::spawn( async move { - let global_guard = match global_guard { + let sandbox_guard = match sandbox_guard { Some(guard) => guard, - None => runtime.lock_global_for_lifecycle(&lifecycle_guard).await, + None => runtime.lock_sandbox_for_lifecycle(&lifecycle_guard).await, }; runtime .compensate_failed_create_inner( &created, lifecycle_guard, - global_guard, + sandbox_guard, original, ) .await @@ -1417,7 +1391,7 @@ impl ComputeRuntime { &self, created: &Sandbox, lifecycle_guard: SandboxLifecycleGuard, - global_guard: tokio::sync::OwnedMutexGuard<()>, + sandbox_guard: LocalMutationGuard, original: Status, ) -> Status { let sandbox_id = created.object_id(); @@ -1460,7 +1434,7 @@ impl ComputeRuntime { }; self.sandbox_index.update_from_sandbox(&transition.deleting); self.sandbox_watch_bus.notify(sandbox_id); - drop(global_guard); + drop(sandbox_guard); let delete_result = self .delete_backend_after_failed_create(sandbox_id, sandbox_name) @@ -1556,7 +1530,7 @@ impl ComputeRuntime { let sandbox_id = candidate.object_id().to_string(); let sandbox_name = candidate.object_name().to_string(); let lifecycle_guard = self.lifecycle_gates.lock_for(&sandbox_id).await; - let global_guard = self.lock_global_for_lifecycle(&lifecycle_guard).await; + let sandbox_guard = self.lock_sandbox_for_lifecycle(&lifecycle_guard).await; let current = self .store .get_message::(&sandbox_id) @@ -1604,7 +1578,7 @@ impl ComputeRuntime { self.sandbox_watch_bus.notify(&sandbox_id); (previous, stopping) }; - drop(global_guard); + drop(sandbox_guard); // Once the durable transition is committed, request cancellation must // not cancel the driver operation and strand the sandbox in @@ -1663,7 +1637,7 @@ impl ComputeRuntime { match result { Ok(_) => { - let _global_guard = self.lock_global_for_lifecycle(&lifecycle_guard).await; + let _sandbox_guard = self.lock_sandbox_for_lifecycle(&lifecycle_guard).await; let latest = self .store .get_message::(&sandbox_id) @@ -1741,7 +1715,7 @@ impl ComputeRuntime { let sandbox_id = candidate.object_id().to_string(); let sandbox_name = candidate.object_name().to_string(); let lifecycle_guard = self.lifecycle_gates.lock_for(&sandbox_id).await; - let global_guard = self.lock_global_for_lifecycle(&lifecycle_guard).await; + let sandbox_guard = self.lock_sandbox_for_lifecycle(&lifecycle_guard).await; let mut current = self .store .get_message::(&sandbox_id) @@ -1887,7 +1861,7 @@ impl ComputeRuntime { } } }; - drop(global_guard); + drop(sandbox_guard); // The durable `Starting` transition commits the operation. Let an // owned worker finish it even if the initiating RPC is canceled. @@ -2004,7 +1978,7 @@ impl ComputeRuntime { .compensate_successful_start(&lifecycle_guard, &starting, &previous, status) .await); } - let global_guard = self.lock_global_for_lifecycle(&lifecycle_guard).await; + let sandbox_guard = self.lock_sandbox_for_lifecycle(&lifecycle_guard).await; let latest = if self.supports_sandbox_authentication() { let driver_name = self.configured_driver_name().to_string(); let persisted = self @@ -2023,7 +1997,7 @@ impl ComputeRuntime { match persisted { Ok(sandbox) => sandbox, Err(error) => { - drop(global_guard); + drop(sandbox_guard); let status = Status::internal(format!( "persist compute runtime identity failed: {error}" )); @@ -2194,7 +2168,7 @@ impl ComputeRuntime { } }; - let _global_guard = self.lock_global_for_lifecycle(lifecycle_guard).await; + let _sandbox_guard = self.lock_sandbox_for_lifecycle(lifecycle_guard).await; if self.restore_lifecycle_snapshot(&settled, previous).await { original } else { @@ -2212,8 +2186,9 @@ impl ComputeRuntime { /// state before deciding whether the pre-operation snapshot is still true. /// /// A transport error can arrive after the runtime applied stop or start. - /// The driver lookup deliberately runs without the process-wide lock; the - /// exact transition resource version then fences the recovery write. + /// The driver lookup deliberately runs without the sandbox's local + /// mutation lock; the exact transition resource version then fences the + /// recovery write. async fn recover_failed_lifecycle( &self, lifecycle_guard: &SandboxLifecycleGuard, @@ -2229,7 +2204,7 @@ impl ComputeRuntime { ) .await .unwrap_or_else(|_| Err("compute lifecycle reconciliation timed out".into())); - let _global_guard = self.lock_global_for_lifecycle(lifecycle_guard).await; + let _sandbox_guard = self.lock_sandbox_for_lifecycle(lifecycle_guard).await; match observed { Ok(Some(snapshot)) if snapshot.id == sandbox_id && snapshot.status.is_some() => { @@ -2446,7 +2421,7 @@ impl ComputeRuntime { target: SandboxDeleteTarget, ) -> Result { let delete_guard = self.lifecycle_gates.lock_for(&target.sandbox_id).await; - let global_guard = self.lock_global_for_lifecycle(&delete_guard).await; + let sandbox_guard = self.lock_sandbox_for_lifecycle(&delete_guard).await; // There is no await between acquiring the initial guards and spawning // the worker. From this commitment point onward, request cancellation @@ -2457,7 +2432,7 @@ impl ComputeRuntime { let request_span = tracing::Span::current(); tokio::spawn( async move { - Box::pin(runtime.delete_sandbox_inner(target, delete_guard, global_guard)).await + Box::pin(runtime.delete_sandbox_inner(target, delete_guard, sandbox_guard)).await } .instrument(request_span), ) @@ -2473,7 +2448,7 @@ impl ComputeRuntime { &self, target: SandboxDeleteTarget, delete_guard: SandboxLifecycleGuard, - guard: tokio::sync::OwnedMutexGuard<()>, + guard: LocalMutationGuard, ) -> Result { let current = self .store @@ -2660,7 +2635,7 @@ impl ComputeRuntime { delete_guard: &SandboxLifecycleGuard, sandbox_id: &str, ) -> bool { - let _guard = self.lock_global_for_lifecycle(delete_guard).await; + let _guard = self.lock_sandbox_for_lifecycle(delete_guard).await; for attempt in 1..=DELETE_PHASE_CAS_RETRY_LIMIT { let record = match self.store.get(Sandbox::object_type(), sandbox_id).await { Ok(Some(record)) => record, @@ -2718,10 +2693,10 @@ impl ComputeRuntime { } /// Removes the sandbox by stable ID only when the expected resource - /// version still owns the row. The caller holds `sync_lock`; a successful - /// delete also removes sandbox-owned records, while successful or - /// already-completed removal clears this replica's index and watch/log - /// buses. + /// version still owns the row. The caller holds this sandbox's local + /// mutation lock; a successful delete also removes sandbox-owned records, + /// while successful or already-completed removal clears this replica's + /// index and watch/log buses. async fn remove_sandbox_record_if_version_locked( &self, sandbox_id: &str, @@ -2806,10 +2781,11 @@ impl ComputeRuntime { /// Resolves an ambiguous driver delete error without overwriting newer /// gateway state. /// - /// The external lookup runs without `sync_lock`. Recovery then uses the - /// exact `Deleting` resource version to apply one of three outcomes: - /// reconcile an observed backend snapshot, remove a confirmed-absent - /// backend, or restore the pre-delete snapshot when lookup is inconclusive. + /// The external lookup runs without the sandbox's local mutation lock. + /// Recovery then uses the exact `Deleting` resource version to apply one of + /// three outcomes: reconcile an observed backend snapshot, remove a + /// confirmed-absent backend, or restore the pre-delete snapshot when lookup + /// is inconclusive. async fn recover_failed_delete( &self, delete_guard: &SandboxLifecycleGuard, @@ -2819,9 +2795,9 @@ impl ComputeRuntime { let sandbox_name = transition.deleting.object_name(); let deleting_resource_version = sandbox_resource_version(&transition.deleting); - // The driver lookup is deliberately outside the process-wide guard. + // The driver lookup is deliberately outside the local mutation lock. let observed = self.get_driver_sandbox(sandbox_id, sandbox_name).await; - let _guard = self.lock_global_for_lifecycle(delete_guard).await; + let _guard = self.lock_sandbox_for_lifecycle(delete_guard).await; match observed { Ok(Some(snapshot)) if snapshot.id == sandbox_id && snapshot.status.is_some() => { @@ -2965,7 +2941,8 @@ impl ComputeRuntime { } } - /// Handles a recovery CAS conflict while the caller holds `sync_lock`. + /// Handles a recovery CAS conflict while the caller holds the sandbox's + /// local mutation lock. /// Another replica may have removed the durable row during the external /// driver lookup; in that case this replica still needs local cleanup. async fn handle_delete_recovery_conflict( @@ -3669,7 +3646,7 @@ impl ComputeRuntime { } async fn mark_sandbox_error(&self, sandbox: &Sandbox, reason: &str, message: &str) { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox.object_id()).await; let sandbox_id = sandbox.object_id().to_string(); let reason = reason.to_string(); let message = message.to_string(); @@ -3715,7 +3692,7 @@ impl ComputeRuntime { /// `Provisioning` with a `Resumed` Ready condition. Returns `true` if the /// store update succeeded. async fn clear_recoverable_error(&self, sandbox: &Sandbox) -> bool { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox.object_id()).await; let sandbox_id = sandbox.object_id().to_string(); match self .store @@ -4000,7 +3977,7 @@ impl ComputeRuntime { } let lifecycle_guard = self.lifecycle_gates.lock_for(sandbox_id).await; - let global_guard = self.lock_global_for_lifecycle(&lifecycle_guard).await; + let sandbox_guard = self.lock_sandbox_for_lifecycle(&lifecycle_guard).await; let Some(current) = self .store .get_message::(sandbox_id) @@ -4081,7 +4058,7 @@ impl ComputeRuntime { }; self.sandbox_index.update_from_sandbox(&claimed); self.sandbox_watch_bus.notify(sandbox_id); - drop(global_guard); + drop(sandbox_guard); let runtime = self.clone(); let owned = claimed.clone(); @@ -4351,7 +4328,7 @@ impl ComputeRuntime { driver_status: Status, ) -> Result<(), String> { let sandbox_id = settled.object_id(); - let _global_guard = self.lock_global_for_lifecycle(lifecycle_guard).await; + let _sandbox_guard = self.lock_sandbox_for_lifecycle(lifecycle_guard).await; let Some(current) = self .store .get_message::(sandbox_id) @@ -4554,7 +4531,7 @@ impl ComputeRuntime { } async fn apply_sandbox_update(&self, mut incoming: DriverSandbox) -> Result<(), String> { - let guard = self.sync_lock.lock().await; + let guard = self.lock_sandbox_local(&incoming.id).await; let mut existing = self .store .get(Sandbox::object_type(), &incoming.id) @@ -4580,9 +4557,9 @@ impl ComputeRuntime { // stop the new generation before the replacement supervisor connects. // The replacement may already be Ready when the old exit arrives, so // terminal container snapshots need the same check after readiness. - // Release the global watch lock, wait for that lifecycle operation, - // and then reread both the driver and store before applying an - // authoritative observation. Taking the per-sandbox gate only for + // Release this sandbox's local mutation lock, wait for that lifecycle + // operation, and then reread both the driver and store before applying + // an authoritative observation. Taking the per-sandbox gate only for // these ambiguous snapshots avoids delaying unrelated watch events behind // slow lifecycle operations. let existing_name = existing_sandbox.as_ref().map_or_else( @@ -4592,7 +4569,7 @@ impl ComputeRuntime { drop(guard); let _lifecycle_guard = self.lifecycle_gates.lock_for(&incoming.id).await; let observed = self.get_driver_sandbox(&incoming.id, &existing_name).await; - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(&incoming.id).await; existing = self .store .get(Sandbox::object_type(), &incoming.id) @@ -4714,7 +4691,7 @@ impl ComputeRuntime { sandbox_id: &str, terminal_delivery_finalized: bool, ) -> Result<(), String> { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox_id).await; // A replacement session may already belong to another gateway. Do not // let cleanup from this replica overwrite the replacement's Ready state. @@ -4794,7 +4771,7 @@ impl ComputeRuntime { instance_id: Option<&str>, terminal_delivery_finalized: bool, ) -> Result<(), String> { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox_id).await; let existing = self .store .get_message::(sandbox_id) @@ -4956,7 +4933,7 @@ impl ComputeRuntime { instance_id: &str, exit_code: i32, ) -> Result<(), String> { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox_id).await; let Some(existing) = self .store .get_message::(sandbox_id) @@ -5050,7 +5027,7 @@ impl ComputeRuntime { sandbox_id: &str, instance_id: &str, ) -> Result<(), String> { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox_id).await; let Some(sandbox) = self .store .get_message::(sandbox_id) @@ -5112,7 +5089,7 @@ impl ComputeRuntime { sandbox_id: &str, instance_id: &str, ) -> Result<(), String> { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox_id).await; let Some(sandbox) = self .store .get_message::(sandbox_id) @@ -5167,7 +5144,7 @@ impl ComputeRuntime { } async fn apply_deleted(&self, sandbox_id: &str) -> Result<(), String> { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(sandbox_id).await; self.apply_deleted_locked(sandbox_id).await } @@ -5270,9 +5247,9 @@ impl ComputeRuntime { /// The gate check is synchronous, but the actual `DeleteSandbox` RPC is /// always deferred to a background task, never awaited inline: both /// call sites run while holding a broader lock (the watch loop's - /// sequential event processing; the prune sweep's gateway-wide - /// `sync_lock`), and a slow or stuck driver call must never block that - /// wider scope. The gate itself is held for the background call's + /// sequential event processing; the prune sweep's local mutation lock + /// for this sandbox), and a slow or stuck driver call must never block + /// that wider scope. The gate itself is held for the background call's /// duration, so this still can't race a concurrent request-side /// operation — only the potentially-slow RPC is backgrounded. fn spawn_driver_sandbox_cleanup(&self, sandbox_id: &str, sandbox_name: &str) { @@ -5442,7 +5419,7 @@ impl ComputeRuntime { delete_guard: &SandboxLifecycleGuard, sandbox_id: &str, ) -> Result { - let _guard = self.lock_global_for_lifecycle(delete_guard).await; + let _guard = self.lock_sandbox_for_lifecycle(delete_guard).await; let record = self .store .get(Sandbox::object_type(), sandbox_id) @@ -5473,7 +5450,7 @@ impl ComputeRuntime { sweep_started_at_ms: i64, ) -> Result<(), String> { let expected_resource_version = { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(&snapshot.id).await; let Some(existing) = self .store .get(Sandbox::object_type(), &snapshot.id) @@ -5496,7 +5473,7 @@ impl ComputeRuntime { return Ok(()); }; - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(&snapshot.id).await; let Some(existing) = self .store .get(Sandbox::object_type(), &snapshot.id) @@ -5522,7 +5499,7 @@ impl ComputeRuntime { grace_ms: i64, ) -> Result<(), String> { let (sandbox_id, sandbox_name, expected_resource_version, age_ms) = { - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(&record.id).await; let Some(current_record) = self .store .get(Sandbox::object_type(), &record.id) @@ -5557,7 +5534,7 @@ impl ComputeRuntime { let current = self.get_driver_sandbox(&sandbox_id, &sandbox_name).await?; - let _guard = self.sync_lock.lock().await; + let _guard = self.lock_sandbox_local(&sandbox_id).await; let Some(current_record) = self .store .get(Sandbox::object_type(), &sandbox_id) @@ -5632,10 +5609,10 @@ impl ComputeRuntime { ); // The driver's own snapshot never reported this sandbox, so no // request-side DeleteSandbox call is coming for it either — release - // driver-owned resources in the background. This function holds - // `sync_lock` (the gateway-wide state guard) through the rest of its - // body, so the driver call must not be awaited here: doing so would - // block every other sandbox operation gateway-wide on a single, + // driver-owned resources in the background. This function holds this + // sandbox's local mutation lock through the rest of its body, so the + // driver call must not be awaited here: doing so would block every + // operation on this sandbox and any global mutation on a single, // potentially slow or stuck driver RPC. self.spawn_driver_sandbox_cleanup(&sandbox_id, &sandbox_name); self.apply_deleted_if_version_locked(&sandbox, expected_resource_version) @@ -7342,7 +7319,7 @@ pub fn new_test_runtime_with_driver( sandbox_watch_bus: SandboxWatchBus::new(), tracing_log_bus: TracingLogBus::new(), supervisor_sessions: Arc::new(SupervisorSessionRegistry::new()), - sync_lock: Arc::new(Mutex::new(())), + mutation_locks: Arc::new(LocalMutationLocks::new()), lifecycle_gates: Arc::new(LifecycleGateRegistry::default()), replica_id: "test-replica".to_string(), rootfs_tar_staging: Arc::new(rootfs_tar::RootfsTarStagingRegistry::disabled()), @@ -8436,7 +8413,7 @@ mod tests { sandbox_watch_bus: SandboxWatchBus::new(), tracing_log_bus: TracingLogBus::new(), supervisor_sessions: Arc::new(SupervisorSessionRegistry::new()), - sync_lock: Arc::new(Mutex::new(())), + mutation_locks: Arc::new(LocalMutationLocks::new()), lifecycle_gates: Arc::new(LifecycleGateRegistry::default()), replica_id: "test-replica".to_string(), rootfs_tar_staging: Arc::new(rootfs_tar::RootfsTarStagingRegistry::disabled()), @@ -8824,13 +8801,15 @@ mod tests { tokio::time::timeout(Duration::from_secs(5), driver.create_started.notified()) .await .unwrap(); - let guard = runtime.sync_lock.clone().lock_owned().await; - let held_references = Arc::strong_count(&runtime.sync_lock); + let guard = runtime.lock_sandbox_local("sb-cancel-cleanup-lock").await; + let sandbox_key = + crate::persistence::MutationLockKey::Sandbox("sb-cancel-cleanup-lock").advisory_key(); + let held_references = runtime.mutation_locks.key_references(sandbox_key); driver.release_create(); - // An additional owned reference shows cleanup has started waiting. - // The cleanup worker must be owned before waiting for this guard. + // An additional reference to the sandbox key shows cleanup has started + // waiting. The cleanup worker must be owned before waiting for this guard. tokio::time::timeout(Duration::from_secs(5), async { - while Arc::strong_count(&runtime.sync_lock) <= held_references { + while runtime.mutation_locks.key_references(sandbox_key) <= held_references { tokio::task::yield_now().await; } }) @@ -10235,6 +10214,66 @@ mod tests { .expect("failed canonical main should delete its ephemeral sandbox"); } + #[tokio::test] + async fn finalized_ephemeral_cleanup_waits_only_for_its_sandbox_lock() { + let driver = ControlledDriver::new(); + let runtime = test_runtime(driver.clone()).await; + let mut sandbox = sandbox_record("sb-1", "sandbox-a", SandboxPhase::Provisioning); + sandbox.metadata.as_mut().unwrap().annotations.insert( + "openshell.nvidia.com/retention".to_string(), + "ephemeral".to_string(), + ); + runtime.store.put_message(&sandbox).await.unwrap(); + runtime + .supervisor_session_connected("sb-1", "instance-1") + .await + .unwrap(); + runtime + .report_main_process_exit("sb-1", "instance-1", 0) + .await + .unwrap(); + runtime + .finalize_main_process_exit("sb-1", "instance-1") + .await + .unwrap(); + + let unrelated = runtime + .mutation_guard(MutationScope::sandbox("default", "sb-2")) + .await + .unwrap(); + let held = runtime.lock_sandbox_local("sb-1").await; + let mut cleanup = tokio::spawn({ + let runtime = runtime.clone(); + async move { + runtime + .cleanup_finalized_ephemeral_sandbox("sb-1", "instance-1") + .await + } + }); + assert!( + tokio::time::timeout(Duration::from_millis(100), &mut cleanup) + .await + .is_err(), + "cleanup should wait for the sandbox's local lock" + ); + assert_eq!(driver.delete_calls(), 0); + + drop(held); + tokio::time::timeout(Duration::from_secs(5), cleanup) + .await + .expect("cleanup should not wait for another sandbox's guard") + .expect("cleanup task") + .unwrap(); + drop(unrelated); + tokio::time::timeout(Duration::from_secs(1), async { + while driver.delete_calls() == 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("cleanup should delete once the sandbox lock is released"); + } + #[tokio::test] async fn conflicting_duplicate_main_process_exit_is_acknowledged() { let runtime = test_runtime(Arc::new(TestDriver::default())).await; @@ -13366,8 +13405,8 @@ mod tests { ); // The driver call is backgrounded (see `spawn_driver_sandbox_cleanup`) - // so the prune sweep never awaits it while holding the gateway-wide - // sync_lock; wait for it to actually land before asserting on it. + // so the prune sweep never awaits it while holding the sandbox's local + // mutation lock; wait for it to actually land before asserting on it. tokio::time::timeout(Duration::from_secs(1), driver.delete_started.notified()) .await .expect("background driver cleanup did not run"); @@ -13380,7 +13419,7 @@ mod tests { #[tokio::test] async fn prune_sweep_does_not_block_on_a_stuck_driver_delete_call() { // Regression test: the prune sweep's driver cleanup must not be - // awaited while holding `sync_lock` (the gateway-wide state guard). + // awaited while holding the sandbox's local mutation lock. // Block the driver's delete call indefinitely and confirm the sweep // itself still completes promptly and removes the store record. let driver = ControlledDriver::new(); @@ -16370,7 +16409,7 @@ mod tests { let mut runtime = test_runtime(driver.clone()).await; enable_runtime_identity_binding(&mut runtime); let mut other = runtime.clone(); - other.sync_lock = Arc::new(Mutex::new(())); + other.mutation_locks = Arc::new(LocalMutationLocks::new()); other.lifecycle_gates = Arc::new(LifecycleGateRegistry::default()); other.driver_info.gateway_manages_lifecycle = true; let mut sandbox = sandbox_record( @@ -16385,7 +16424,7 @@ mod tests { let create = tokio::spawn(async move { creating.create_sandbox(sandbox, None, false).await }); driver.create_started.notified().await; - let held = runtime.sync_lock.lock().await; + let held = runtime.lock_sandbox_local("sb-result-owner").await; driver.release_create(); wait_driver_pending(&other, "sb-result-owner", false).await; driver.set_runtime_identity("recovery-runtime"); @@ -16468,7 +16507,7 @@ mod tests { driver.block_stop(); let runtime = test_runtime(driver.clone()).await; let mut other = runtime.clone(); - other.sync_lock = Arc::new(Mutex::new(())); + other.mutation_locks = Arc::new(LocalMutationLocks::new()); other.lifecycle_gates = Arc::new(LifecycleGateRegistry::default()); let now = openshell_core::time::now_ms(); let mut sandbox = sandbox_record("sb-stop-lease", "stop-lease", SandboxPhase::Provisioning); @@ -16557,7 +16596,7 @@ mod tests { let mut runtime = test_runtime(driver.clone()).await; enable_runtime_identity_binding(&mut runtime); let mut other = runtime.clone(); - other.sync_lock = Arc::new(Mutex::new(())); + other.mutation_locks = Arc::new(LocalMutationLocks::new()); other.lifecycle_gates = Arc::new(LifecycleGateRegistry::default()); let mut sandbox = sandbox_record("sb-restart-pending", "restart-pending", SandboxPhase::Ready); @@ -16661,7 +16700,7 @@ mod tests { driver.block_start(); let runtime = test_runtime(driver.clone()).await; let mut other = runtime.clone(); - other.sync_lock = Arc::new(Mutex::new(())); + other.mutation_locks = Arc::new(LocalMutationLocks::new()); other.lifecycle_gates = Arc::new(LifecycleGateRegistry::default()); let mut sandbox = sandbox_record("sb-auto-owned", "auto-owned", SandboxPhase::Ready); sandbox @@ -16856,12 +16895,12 @@ mod tests { owned.clone() }; let gate = runtime.lifecycle_gates.lock_for(sandbox.object_id()).await; - let global = runtime.lock_global_for_lifecycle(&gate).await; + let sandbox_guard = runtime.lock_sandbox_for_lifecycle(&gate).await; let error = runtime .compensate_failed_create( &owned, gate, - Some(global), + Some(sandbox_guard), Status::internal("missing binding"), ) .await; @@ -16972,7 +17011,7 @@ mod tests { let create = tokio::spawn(async move { creating.create_sandbox(sandbox, None, false).await }); driver.create_started.notified().await; - let held = runtime.sync_lock.lock().await; + let held = runtime.lock_sandbox_local("sb-monitor").await; sqlx::query("ALTER TABLE objects RENAME TO temporarily_hidden_objects") .execute(&pool) .await @@ -17047,7 +17086,7 @@ mod tests { ); let mut restarted = runtime.clone(); restarted.lifecycle_gates = Arc::new(LifecycleGateRegistry::default()); - restarted.sync_lock = Arc::new(Mutex::new(())); + restarted.mutation_locks = Arc::new(LocalMutationLocks::new()); restarted.driver_info.gateway_manages_lifecycle = true; restarted .start_persisted_sandboxes_with_authentication( @@ -17385,7 +17424,7 @@ mod tests { .reconcile_provisioning_deadlines(now + 1_000) .await .unwrap(); - let _global_guard = runtime.lock_global_for_lifecycle(&lifecycle_guard).await; + let _sandbox_guard = runtime.lock_sandbox_for_lifecycle(&lifecycle_guard).await; let result = runtime .begin_sandbox_delete_with_initial_snapshot( sandbox.object_id(), @@ -17412,7 +17451,7 @@ mod tests { driver.track_compute.store(true, Ordering::SeqCst); let runtime = test_runtime(driver.clone()).await; let mut other = runtime.clone(); - other.sync_lock = Arc::new(Mutex::new(())); + other.mutation_locks = Arc::new(LocalMutationLocks::new()); other.lifecycle_gates = Arc::new(LifecycleGateRegistry::default()); other.replica_id = "other-replica".into(); let now = openshell_core::time::now_ms(); @@ -17480,7 +17519,7 @@ mod tests { // STOP has observed absence. Keep its final reread blocked while the // original CREATE succeeds and durably clears pending ownership. - let held = other.sync_lock.lock().await; + let held = other.lock_sandbox_local("sb-create-retry").await; driver.release_stop(); driver.stop_finished.notified().await; driver.release_create(); @@ -17869,7 +17908,7 @@ mod tests { .unwrap() .unwrap(); let gate = runtime.lifecycle_gates.lock_for("sb-prepare").await; - let global = runtime.lock_global_for_lifecycle(&gate).await; + let sandbox_guard = runtime.lock_sandbox_for_lifecycle(&gate).await; assert!( runtime .claim_provisioning_timeout(&sandbox, 600_000) @@ -17902,7 +17941,7 @@ mod tests { assert!(record.admission_start_time.is_none()); assert!(record.preparation_deadline.is_some()); assert!(record.cleanup_completed_time.is_none()); - drop(global); + drop(sandbox_guard); runtime .reclaim_provisioning_timeout(&expired, &gate) .await @@ -17933,7 +17972,7 @@ mod tests { let runtime = test_runtime(driver.clone()).await; let sandbox = seed_provisioning_attempt(&runtime).await; let gate = runtime.lifecycle_gates.lock_for("sb-ttl").await; - let global = runtime.lock_global_for_lifecycle(&gate).await; + let sandbox_guard = runtime.lock_sandbox_for_lifecycle(&gate).await; assert!( runtime .claim_provisioning_timeout(&sandbox, 299_999) @@ -17946,7 +17985,7 @@ mod tests { .await .unwrap() .unwrap(); - drop(global); + drop(sandbox_guard); assert_eq!(expired.phase(), i32::from(SandboxPhase::Error)); assert!( ready_condition(&expired) @@ -18040,13 +18079,13 @@ mod tests { let runtime = test_runtime(driver.clone()).await; let sandbox = seed_provisioning_attempt(&runtime).await; let gate = runtime.lifecycle_gates.lock_for("sb-ttl").await; - let global = runtime.lock_global_for_lifecycle(&gate).await; + let sandbox_guard = runtime.lock_sandbox_for_lifecycle(&gate).await; let expired = runtime .claim_provisioning_timeout(&sandbox, 300_000) .await .unwrap() .unwrap(); - drop(global); + drop(sandbox_guard); runtime .reclaim_provisioning_timeout(&expired, &gate) .await @@ -18157,6 +18196,43 @@ mod tests { assert_eq!(driver.stop_calls(), 1); } + #[tokio::test] + async fn provisioning_reconcile_waits_for_a_local_workspace_writer() { + let runtime = test_runtime(ControlledDriver::new()).await; + runtime + .store + .put_message(&sandbox_record( + "sb-ws-writer", + "ws-writer", + SandboxPhase::Provisioning, + )) + .await + .unwrap(); + // Reconcile re-derives configuration from provider and profile + // records, so it waits for their writers, which hold X(workspace). + let writer = runtime + .mutation_guard(MutationScope::Workspace("default")) + .await + .unwrap(); + let mut reconcile = { + let runtime = runtime.clone(); + let now = openshell_core::time::now_ms(); + tokio::spawn(async move { runtime.reconcile_provisioning_deadlines(now).await }) + }; + assert!( + tokio::time::timeout(Duration::from_millis(100), &mut reconcile) + .await + .is_err(), + "provisioning reconcile must wait for the workspace writer" + ); + drop(writer); + tokio::time::timeout(Duration::from_secs(5), reconcile) + .await + .expect("provisioning reconcile proceeds once the writer releases") + .unwrap() + .unwrap(); + } + #[tokio::test] async fn provisioning_worker_uses_policy_commit_time_not_scan_time() { use crate::policy_store::PolicyStoreExt; @@ -18294,13 +18370,13 @@ mod tests { let runtime = test_runtime(driver.clone()).await; let sandbox = seed_provisioning_attempt(&runtime).await; let gate = runtime.lifecycle_gates.lock_for("sb-ttl").await; - let global = runtime.lock_global_for_lifecycle(&gate).await; + let sandbox_guard = runtime.lock_sandbox_for_lifecycle(&gate).await; let expired = runtime .claim_provisioning_timeout(&sandbox, 300_000) .await .unwrap() .unwrap(); - drop(global); + drop(sandbox_guard); runtime .reclaim_provisioning_timeout(&expired, &gate) .await diff --git a/crates/openshell-server/src/compute/mutation_guard.rs b/crates/openshell-server/src/compute/mutation_guard.rs new file mode 100644 index 0000000000..931035031e --- /dev/null +++ b/crates/openshell-server/src/compute/mutation_guard.rs @@ -0,0 +1,1257 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Hierarchical mutation guards: a process-local keyed lock table plus, on +//! `PostgreSQL`, the matching advisory locks on the dedicated lock pool. +//! +//! See [`crate::persistence::mutation_lock`] for the key hierarchy and the +//! ordering rules. In short: lifecycle gate first, then local keys in +//! ascending order, then `PostgreSQL` keys in ascending order on one +//! connection. A task never acquires a guard while it holds one. + +use super::{ComputeRuntime, SandboxLifecycleGuard}; +use crate::gateway_metrics::{self, LockScope}; +use crate::grpc::workspace::DEFAULT_WORKSPACE_NAME; +use crate::persistence::mutation_lock::MUTATION_LOCK_TIMEOUT; +use crate::persistence::{ + DistributedMutationGuard, LockMode, MutationLockKey, MutationLockSet, PersistenceError, + PersistenceResult, +}; +use openshell_core::ObjectWorkspace; +use openshell_core::proto::Sandbox; +use std::collections::HashMap; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Mutex as StdMutex, Weak}; +use std::time::Duration; +use tokio::sync::{OwnedRwLockReadGuard, OwnedRwLockWriteGuard, RwLock}; +use tracing::warn; + +/// What a guarded mutation reads and writes, which selects its lock set. +/// +/// Every mutation whose invariant spans several persisted objects must take +/// the narrowest scope that still conflicts with every writer of the objects +/// it validates against. Global settings and policy writers and +/// platform-scope profile writers hold the global key exclusively; provider +/// and workspace-scoped profile writers hold their workspace key exclusively; +/// sandbox-scoped settings and policy writers hold only their sandbox key +/// exclusively. A new cross-object writer needs a scope from the same table. +#[derive(Clone, Copy, Debug)] +pub enum MutationScope<'a> { + /// Global policy/settings and platform-scope profiles. Excludes every + /// other scope fleet-wide. + Global, + /// Provider and workspace-scoped profile mutations. `""` (platform) + /// behaves as `Global`. + Workspace(&'a str), + /// Any mutation of one sandbox's records, admin or supervisor. + Sandbox { + workspace: &'a str, + sandbox_id: &'a str, + }, +} + +impl<'a> MutationScope<'a> { + pub(crate) const fn sandbox(workspace: &'a str, sandbox_id: &'a str) -> Self { + Self::Sandbox { + workspace, + sandbox_id, + } + } + + /// Keys and modes of this scope: + /// + /// - `Global` and `Workspace("")`: X(global). + /// - `Workspace(ws)`: S(global) X(workspace). + /// - `Sandbox`: S(global) S(workspace) X(sandbox). A legacy sandbox with + /// an empty workspace locks the default workspace, where its providers + /// resolve. + pub(crate) fn lock_set(&self) -> MutationLockSet { + let mut set = MutationLockSet::default(); + match *self { + Self::Global => set.insert(MutationLockKey::Global, LockMode::Exclusive), + Self::Workspace("") => { + set.insert(MutationLockKey::Global, LockMode::Exclusive); + } + Self::Workspace(workspace) => { + set.insert(MutationLockKey::Global, LockMode::Shared); + set.insert(MutationLockKey::Workspace(workspace), LockMode::Exclusive); + } + Self::Sandbox { + workspace, + sandbox_id, + } => { + set.insert(MutationLockKey::Global, LockMode::Shared); + set.insert( + MutationLockKey::Workspace(sandbox_workspace_key(workspace)), + LockMode::Shared, + ); + set.insert(MutationLockKey::Sandbox(sandbox_id), LockMode::Exclusive); + } + } + set + } + + /// The metric label of this scope. + pub(crate) const fn lock_scope(&self) -> LockScope { + match self { + Self::Global => LockScope::Global, + Self::Workspace(workspace) if workspace.is_empty() => LockScope::Global, + Self::Workspace(_) => LockScope::Workspace, + Self::Sandbox { .. } => LockScope::Sandbox, + } + } +} + +/// Workspace key of a sandbox scope. Legacy sandboxes may carry an empty +/// workspace; their providers resolve in the default workspace. +fn sandbox_workspace_key(workspace: &str) -> &str { + if workspace.is_empty() { + DEFAULT_WORKSPACE_NAME + } else { + workspace + } +} + +/// Process-local table of mutation lock keys. +/// +/// Entries are weak, so a key disappears once no guard holds it and no +/// acquisition waits on it. Tokio's `RwLock` is fair and write-preferring: a +/// queued exclusive request blocks later shared requests on the same key. +#[derive(Debug)] +pub struct LocalMutationLocks { + entries: StdMutex>>>, + timeout_ms: AtomicU64, +} + +impl LocalMutationLocks { + pub(crate) fn new() -> Self { + Self { + entries: StdMutex::new(HashMap::new()), + timeout_ms: AtomicU64::new(duration_millis(MUTATION_LOCK_TIMEOUT)), + } + } + + fn lock_for(&self, key: i64) -> Arc> { + let mut entries = self + .entries + .lock() + .expect("mutation lock registry lock poisoned"); + entries.retain(|_, lock| lock.strong_count() > 0); + + if let Some(lock) = entries.get(&key).and_then(Weak::upgrade) { + return lock; + } + + let lock = Arc::new(RwLock::new(())); + entries.insert(key, Arc::downgrade(&lock)); + lock + } + + /// Acquire `set` in ascending key order. Dropping the future releases the + /// keys already taken and leaves no queue entry behind. + async fn acquire(&self, set: &MutationLockSet) -> LocalMutationGuard { + let mut guards = Vec::new(); + for (key, mode) in set.iter() { + let lock = self.lock_for(key); + guards.push(match mode { + LockMode::Shared => LocalKeyGuard::Shared { + _guard: lock.read_owned().await, + }, + LockMode::Exclusive => LocalKeyGuard::Exclusive { + _guard: lock.write_owned().await, + }, + }); + } + LocalMutationGuard { _guards: guards } + } + + fn timeout(&self) -> Duration { + Duration::from_millis(self.timeout_ms.load(Ordering::Relaxed)) + } + + #[cfg(test)] + pub(crate) fn set_timeout_for_tests(&self, timeout: Duration) { + self.timeout_ms + .store(duration_millis(timeout), Ordering::Relaxed); + } + + /// Entries in the table, including released keys that `lock_for` has not + /// pruned yet. + #[cfg(test)] + pub(crate) fn entry_count(&self) -> usize { + self.entries + .lock() + .expect("mutation lock registry lock poisoned") + .len() + } + + /// Keys still held by a guard or awaited by an acquisition. Does not + /// prune, so it cannot hide a leak in `lock_for`. + #[cfg(test)] + pub(crate) fn live_entry_count(&self) -> usize { + self.entries + .lock() + .expect("mutation lock registry lock poisoned") + .values() + .filter(|lock| lock.strong_count() > 0) + .count() + } + + /// References to `key`'s lock: one per guard holding it and one per + /// acquisition waiting for it. + #[cfg(test)] + pub(crate) fn key_references(&self, key: i64) -> usize { + self.entries + .lock() + .expect("mutation lock registry lock poisoned") + .get(&key) + .map_or(0, Weak::strong_count) + } +} + +fn duration_millis(duration: Duration) -> u64 { + u64::try_from(duration.as_millis()).unwrap_or(u64::MAX) +} + +enum LocalKeyGuard { + Shared { _guard: OwnedRwLockReadGuard<()> }, + Exclusive { _guard: OwnedRwLockWriteGuard<()> }, +} + +/// Process-local keys of one mutation or lifecycle operation. +#[must_use = "dropping the guard releases the local mutation locks"] +pub struct LocalMutationGuard { + _guards: Vec, +} + +/// Local and, on `PostgreSQL`, distributed keys of one guarded mutation. +#[must_use = "dropping the guard releases the mutation locks"] +pub struct MutationGuard { + // Field order is drop order: the database guard goes first so its + // connection returns to the lock pool (and is unlocked) as early as + // possible. + _distributed: DistributedMutationGuard, + _local: LocalMutationGuard, +} + +impl ComputeRuntime { + /// Serialize a cross-object mutation against every conflicting mutation + /// on this and, on `PostgreSQL`, every other replica. + /// + /// One `MUTATION_LOCK_TIMEOUT` deadline covers the local keys, the + /// lock-pool connection, and the advisory locks. Missing it fails with + /// [`PersistenceError::LockTimeout`]: the mutation lock was not acquired, + /// so the caller's guarded writes did not run. A lock connection that + /// `PostgreSQL` does not open with at least `LOCK_CONNECTION_MIN_BUDGET` + /// left fails with [`PersistenceError::Database`]. + pub(crate) async fn mutation_guard( + &self, + scope: MutationScope<'_>, + ) -> PersistenceResult { + let started = tokio::time::Instant::now(); + let deadline = started + self.mutation_locks.timeout(); + let result = self + .acquire_mutation_guard(&scope.lock_set(), deadline) + .await; + match &result { + Ok(_) => gateway_metrics::record_lock_wait(scope.lock_scope(), started.elapsed()), + Err(PersistenceError::LockTimeout(detail)) => { + gateway_metrics::record_lock_timeout(scope.lock_scope()); + warn!( + scope = scope.lock_scope().label(), + waited_ms = duration_millis(started.elapsed()), + detail = %detail, + "mutation lock acquisition timed out" + ); + } + Err(error) => warn!( + scope = scope.lock_scope().label(), + waited_ms = duration_millis(started.elapsed()), + error = %error, + "mutation lock acquisition failed" + ), + } + result + } + + async fn acquire_mutation_guard( + &self, + set: &MutationLockSet, + deadline: tokio::time::Instant, + ) -> PersistenceResult { + let local = tokio::time::timeout_at(deadline, self.mutation_locks.acquire(set)) + .await + .map_err(|_| { + PersistenceError::LockTimeout("waiting for a local mutation lock".into()) + })?; + let distributed = self + .store + .acquire_distributed_mutation_guard(set, deadline) + .await?; + Ok(MutationGuard { + _distributed: distributed, + _local: local, + }) + } + + /// Sandbox-scoped guard for paths that know only the sandbox id, such as + /// supervisor reports. + /// + /// A sandbox's workspace never changes, so one read before locking + /// derives the key set. Callers must re-read the sandbox after locking and + /// never validate against this read. Returns `Ok(None)` when the sandbox + /// does not exist. + pub(crate) async fn sandbox_mutation_guard_by_id( + &self, + sandbox_id: &str, + ) -> PersistenceResult> { + let Some(sandbox) = self.store.get_message::(sandbox_id).await? else { + return Ok(None); + }; + self.mutation_guard(MutationScope::sandbox( + sandbox.object_workspace(), + sandbox_id, + )) + .await + .map(Some) + } + + /// Lifecycle gate, then the sandbox-scoped mutation guard, for a new + /// sandbox. + pub(crate) async fn sandbox_create_guards( + &self, + workspace: &str, + sandbox_id: &str, + ) -> PersistenceResult<(SandboxLifecycleGuard, MutationGuard)> { + let lifecycle_guard = self.lifecycle_gates.lock_for(sandbox_id).await; + let mutation_guard = self + .mutation_guard(MutationScope::sandbox(workspace, sandbox_id)) + .await?; + Ok((lifecycle_guard, mutation_guard)) + } + + /// Local S(global) X(sandbox) for code that already holds the sandbox's + /// lifecycle gate. The guard parameter documents and enforces the + /// lifecycle-gate -> mutation-lock order. + pub(super) async fn lock_sandbox_for_lifecycle( + &self, + lifecycle_guard: &SandboxLifecycleGuard, + ) -> LocalMutationGuard { + self.lock_sandbox_local(&lifecycle_guard.sandbox_id).await + } + + /// Local S(global) X(sandbox) for lifecycle, driver-watch, and reconcile + /// paths. They write only this sandbox and its owned records and rely on + /// compare-and-swap across replicas, so they take no database lock and + /// never exclude another sandbox or a provider writer. + pub(super) async fn lock_sandbox_local(&self, sandbox_id: &str) -> LocalMutationGuard { + self.mutation_locks + .acquire(&MutationLockSet::sandbox_lifecycle(sandbox_id)) + .await + } + + /// Local S(global) S(workspace) X(sandbox), for provisioning-deadline + /// reconciliation, which re-derives configuration from provider and + /// profile records and must not interleave with their local writers. + pub(super) async fn lock_sandbox_local_in_workspace( + &self, + workspace: &str, + sandbox_id: &str, + ) -> LocalMutationGuard { + self.mutation_locks + .acquire(&MutationScope::sandbox(workspace, sandbox_id).lock_set()) + .await + } + + /// Shorten the mutation lock deadline of this runtime and its clones. + #[cfg(test)] + pub(crate) fn set_mutation_lock_timeout_for_tests(&self, timeout: Duration) { + self.mutation_locks.set_timeout_for_tests(timeout); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::gateway_metrics::MetricsCapture; + use crate::persistence::Store; + use crate::persistence::mutation_lock::GLOBAL_MUTATION_LOCK_KEY; + use crate::persistence::test_postgres::TestSchema; + use openshell_core::GetResourceVersion; + use openshell_core::proto::SandboxPhase; + use rand::rngs::StdRng; + use rand::{Rng, SeedableRng}; + use std::sync::atomic::{AtomicBool, AtomicIsize}; + use tokio::task::JoinHandle; + use uuid::Uuid; + + const BLOCKED_FOR: Duration = Duration::from_millis(100); + const PROCEEDS_WITHIN: Duration = Duration::from_secs(5); + + async fn test_runtime() -> ComputeRuntime { + let store = Arc::new( + Store::connect("sqlite::memory:?cache=shared") + .await + .expect("in-memory store"), + ); + super::super::new_test_runtime_for_driver(store, "test").await + } + + fn spawn_guard( + runtime: &ComputeRuntime, + scope: MutationScope<'static>, + ) -> JoinHandle { + let runtime = runtime.clone(); + tokio::spawn(async move { + runtime + .mutation_guard(scope) + .await + .expect("mutation guard acquired") + }) + } + + fn spawn_local(runtime: &ComputeRuntime, sandbox_id: &'static str) -> JoinHandle<()> { + let runtime = runtime.clone(); + tokio::spawn(async move { + drop(runtime.lock_sandbox_local(sandbox_id).await); + }) + } + + async fn assert_blocked(handle: &mut JoinHandle, what: &str) { + assert!( + tokio::time::timeout(BLOCKED_FOR, handle).await.is_err(), + "{what} should wait" + ); + } + + async fn assert_proceeds(handle: JoinHandle, what: &str) -> T { + tokio::time::timeout(PROCEEDS_WITHIN, handle) + .await + .unwrap_or_else(|_| panic!("{what} should proceed")) + .expect("guard task") + } + + fn keys(entries: &[(MutationLockKey<'_>, LockMode)]) -> Vec<(i64, LockMode)> { + let mut keys: Vec<_> = entries + .iter() + .map(|(key, mode)| (key.advisory_key(), *mode)) + .collect(); + keys.sort_unstable(); + keys + } + + #[test] + fn scope_lock_sets_follow_the_hierarchy() { + use LockMode::{Exclusive, Shared}; + use MutationLockKey::{Global, Sandbox as SandboxKey, Workspace}; + + let cases = [ + (MutationScope::Global, keys(&[(Global, Exclusive)])), + (MutationScope::Workspace(""), keys(&[(Global, Exclusive)])), + ( + MutationScope::Workspace("team-a"), + keys(&[(Global, Shared), (Workspace("team-a"), Exclusive)]), + ), + ( + MutationScope::sandbox("team-a", "sb-1"), + keys(&[ + (Global, Shared), + (Workspace("team-a"), Shared), + (SandboxKey("sb-1"), Exclusive), + ]), + ), + ( + MutationScope::sandbox("", "sb-1"), + keys(&[ + (Global, Shared), + (Workspace("default"), Shared), + (SandboxKey("sb-1"), Exclusive), + ]), + ), + ]; + for (scope, expected) in cases { + assert_eq!( + scope.lock_set().iter().collect::>(), + expected, + "{scope:?}" + ); + } + assert_eq!( + MutationLockSet::sandbox_lifecycle("sb-1") + .iter() + .collect::>(), + keys(&[(Global, Shared), (SandboxKey("sb-1"), Exclusive)]) + ); + assert!( + MutationScope::Global + .lock_set() + .iter() + .eq([(GLOBAL_MUTATION_LOCK_KEY, Exclusive)]) + ); + } + + #[test] + fn scope_labels_map_platform_workspace_to_global() { + assert_eq!(MutationScope::Global.lock_scope(), LockScope::Global); + assert_eq!(MutationScope::Workspace("").lock_scope(), LockScope::Global); + assert_eq!( + MutationScope::Workspace("team-a").lock_scope(), + LockScope::Workspace + ); + assert_eq!( + MutationScope::sandbox("", "sb-1").lock_scope(), + LockScope::Sandbox + ); + } + + #[tokio::test] + async fn sandbox_guards_for_different_sandboxes_proceed_concurrently() { + let runtime = test_runtime().await; + let held = runtime + .mutation_guard(MutationScope::sandbox("w1", "a")) + .await + .unwrap(); + + let other = spawn_guard(&runtime, MutationScope::sandbox("w1", "b")); + drop(assert_proceeds(other, "a different sandbox in the same workspace").await); + drop(held); + } + + #[tokio::test] + async fn same_sandbox_guard_waits_until_release() { + let runtime = test_runtime().await; + let held = runtime + .mutation_guard(MutationScope::sandbox("w1", "a")) + .await + .unwrap(); + + let mut same = spawn_guard(&runtime, MutationScope::sandbox("w1", "a")); + assert_blocked(&mut same, "the same sandbox").await; + drop(held); + drop(assert_proceeds(same, "the same sandbox after release").await); + } + + #[tokio::test] + async fn workspace_guard_blocks_sandbox_scope_in_that_workspace_only() { + let runtime = test_runtime().await; + let held = runtime + .mutation_guard(MutationScope::Workspace("w1")) + .await + .unwrap(); + + let mut same_workspace = spawn_guard(&runtime, MutationScope::sandbox("w1", "a")); + assert_blocked(&mut same_workspace, "a sandbox in the held workspace").await; + let other_sandbox = spawn_guard(&runtime, MutationScope::sandbox("w2", "b")); + drop(assert_proceeds(other_sandbox, "a sandbox in another workspace").await); + let other_workspace = spawn_guard(&runtime, MutationScope::Workspace("w2")); + drop(assert_proceeds(other_workspace, "another workspace").await); + + drop(held); + drop(assert_proceeds(same_workspace, "the sandbox after release").await); + } + + #[tokio::test] + async fn global_guard_blocks_every_scope_and_lifecycle_lock() { + let runtime = test_runtime().await; + let held = runtime.mutation_guard(MutationScope::Global).await.unwrap(); + + let mut waiting_guards = vec![ + spawn_guard(&runtime, MutationScope::Global), + spawn_guard(&runtime, MutationScope::Workspace("")), + spawn_guard(&runtime, MutationScope::Workspace("w1")), + spawn_guard(&runtime, MutationScope::sandbox("w1", "a")), + ]; + for waiting in &mut waiting_guards { + assert_blocked(waiting, "a guard behind the global guard").await; + } + let mut lifecycle = spawn_local(&runtime, "b"); + assert_blocked(&mut lifecycle, "a lifecycle lock behind the global guard").await; + let reconcile_runtime = runtime.clone(); + let mut reconcile = tokio::spawn(async move { + drop( + reconcile_runtime + .lock_sandbox_local_in_workspace("w1", "c") + .await, + ); + }); + assert_blocked(&mut reconcile, "a reconcile lock behind the global guard").await; + + drop(held); + for waiting in waiting_guards { + drop(assert_proceeds(waiting, "a guard after release").await); + } + assert_proceeds(lifecycle, "the lifecycle lock after release").await; + assert_proceeds(reconcile, "the reconcile lock after release").await; + } + + #[tokio::test] + async fn queued_global_guard_is_not_starved() { + let runtime = test_runtime().await; + let held = runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + .unwrap(); + + let mut global = spawn_guard(&runtime, MutationScope::Global); + assert_blocked(&mut global, "the global guard behind a sandbox guard").await; + let mut later = spawn_guard(&runtime, MutationScope::sandbox("w", "b")); + assert_blocked(&mut later, "a sandbox guard queued behind the global guard").await; + + drop(held); + let global = assert_proceeds(global, "the queued global guard").await; + assert_blocked(&mut later, "a sandbox guard while the global guard holds").await; + drop(global); + drop(assert_proceeds(later, "the later sandbox guard").await); + } + + #[tokio::test] + async fn lifecycle_lock_excludes_same_sandbox_only() { + let runtime = test_runtime().await; + let held = runtime.lock_sandbox_local("a").await; + + let mut same = spawn_guard(&runtime, MutationScope::sandbox("w", "a")); + assert_blocked(&mut same, "the sandbox held by a lifecycle lock").await; + let other = spawn_guard(&runtime, MutationScope::sandbox("w", "b")); + drop(assert_proceeds(other, "another sandbox").await); + let provider = spawn_guard(&runtime, MutationScope::Workspace("w")); + drop(assert_proceeds(provider, "a provider writer").await); + + drop(held); + drop(assert_proceeds(same, "the sandbox after release").await); + } + + #[tokio::test] + async fn legacy_empty_workspace_sandbox_conflicts_with_default_workspace_writer() { + let runtime = test_runtime().await; + let held = runtime + .mutation_guard(MutationScope::Workspace("default")) + .await + .unwrap(); + + let mut legacy = spawn_guard(&runtime, MutationScope::sandbox("", "a")); + assert_blocked( + &mut legacy, + "a legacy sandbox behind a default-workspace writer", + ) + .await; + drop(held); + drop(assert_proceeds(legacy, "the legacy sandbox after release").await); + } + + #[tokio::test] + async fn local_registry_drops_released_entries() { + let runtime = test_runtime().await; + let sandbox = runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + .unwrap(); + let lifecycle = runtime.lock_sandbox_local("b").await; + assert_eq!(runtime.mutation_locks.entry_count(), 4); + + drop(sandbox); + drop(lifecycle); + assert_eq!(runtime.mutation_locks.live_entry_count(), 0); + + // The next acquisition prunes the released keys, so the table holds + // only the new guard's global and sandbox keys. + let fresh = runtime.lock_sandbox_local("c").await; + assert_eq!(runtime.mutation_locks.entry_count(), 2); + drop(fresh); + } + + #[tokio::test] + async fn cancelled_acquisition_leaves_no_queue_entry() { + let runtime = test_runtime().await; + let held = runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + .unwrap(); + let mut waiter = spawn_guard(&runtime, MutationScope::sandbox("w", "a")); + assert_blocked(&mut waiter, "the waiter").await; + waiter.abort(); + let Err(error) = waiter.await else { + panic!("the aborted waiter should not acquire the guard"); + }; + assert!(error.is_cancelled()); + drop(held); + + let next = spawn_guard(&runtime, MutationScope::sandbox("w", "a")); + drop(assert_proceeds(next, "a new acquisition after the cancelled one").await); + assert_eq!(runtime.mutation_locks.live_entry_count(), 0); + } + + #[tokio::test] + async fn local_timeout_returns_lock_timeout_and_unavailable() { + let runtime = test_runtime().await; + runtime.set_mutation_lock_timeout_for_tests(Duration::from_millis(50)); + let held = runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + .unwrap(); + + let Err(error) = runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + else { + panic!("the second acquisition should time out"); + }; + assert!( + matches!(error, PersistenceError::LockTimeout(_)), + "{error:?}" + ); + let status = crate::grpc::persistence_error_to_status(error, "op"); + assert_eq!(status.code(), tonic::Code::Unavailable); + let details = openshell_core::rpc_error::decode_details(&status).expect("error details"); + assert_eq!( + details.error_info().expect("error info").reason, + "MUTATION_LOCK_TIMEOUT" + ); + drop(held); + assert_eq!(runtime.mutation_locks.live_entry_count(), 0); + } + + #[tokio::test] + async fn lock_metrics_record_wait_and_timeout() { + const WAITS: &str = "openshell_server_mutation_lock_wait_seconds_count{scope=\"sandbox\"}"; + const TIMEOUTS: &str = "openshell_server_mutation_lock_timeouts_total{scope=\"sandbox\"}"; + let metrics = MetricsCapture::install(); + let runtime = test_runtime().await; + runtime.set_mutation_lock_timeout_for_tests(Duration::from_millis(50)); + + let held = runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + .unwrap(); + assert_eq!(metrics.value(WAITS), Some(1)); + assert_eq!(metrics.value(TIMEOUTS), None); + + assert!( + runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + .is_err() + ); + assert_eq!(metrics.value(TIMEOUTS), Some(1)); + assert_eq!(metrics.value(WAITS), Some(1)); + + drop(runtime.lock_sandbox_local("b").await); + assert_eq!(metrics.value(WAITS), Some(1)); + drop(held); + } + + #[tokio::test] + async fn lock_connection_failure_is_not_counted_as_a_timeout() { + const TIMEOUTS: &str = "openshell_server_mutation_lock_timeouts_total{scope=\"sandbox\"}"; + let metrics = MetricsCapture::install(); + let url = crate::persistence::PostgresStore::refusing_url_for_tests().await; + let store = crate::persistence::PostgresStore::connect_lazy_for_tests(&url, 1); + let runtime = + super::super::new_test_runtime_for_driver(Arc::new(Store::Postgres(store)), "test") + .await; + // Leave enough time to open a lock connection, so the refusal is a + // database error. + runtime.set_mutation_lock_timeout_for_tests( + crate::persistence::mutation_lock::LOCK_CONNECTION_MIN_BUDGET + + Duration::from_millis(300), + ); + + match runtime + .mutation_guard(MutationScope::sandbox("w", "a")) + .await + { + Err(PersistenceError::Database(detail)) => assert!( + detail.starts_with("could not open a mutation lock connection"), + "{detail}" + ), + Err(error) => panic!("expected a database error, got {error:?}"), + Ok(_) => panic!("nothing listens, yet the guard was acquired"), + } + assert_eq!(metrics.value(TIMEOUTS), None); + } + + /// Workspaces and sandboxes the random scope mix draws from. + const MIX_WORKSPACES: usize = 3; + const MIX_SANDBOXES: usize = 8; + + /// Names of the random scope mix's workspaces and sandboxes. + struct MixNames { + workspaces: [String; MIX_WORKSPACES], + sandboxes: [String; MIX_SANDBOXES], + } + + impl MixNames { + fn new(prefix: &str) -> Self { + Self { + workspaces: std::array::from_fn(|index| format!("{prefix}w{index}")), + sandboxes: std::array::from_fn(|index| format!("{prefix}s{index}")), + } + } + } + + /// One step of a random-mix task, as indices into [`MixNames`]. + #[derive(Clone, Copy)] + enum MixOp { + Global, + Workspace(usize), + Sandbox(usize, usize), + Lifecycle(usize), + GatedLifecycle(usize), + } + + impl MixOp { + fn random(rng: &mut StdRng) -> Self { + match rng.random_range(0..4) { + 0 => Self::Global, + 1 => Self::Workspace(rng.random_range(0..MIX_WORKSPACES)), + 2 => Self::Sandbox( + rng.random_range(0..MIX_WORKSPACES), + rng.random_range(0..MIX_SANDBOXES), + ), + _ => { + let sandbox = rng.random_range(0..MIX_SANDBOXES); + if rng.random_bool(0.5) { + Self::GatedLifecycle(sandbox) + } else { + Self::Lifecycle(sandbox) + } + } + } + } + + /// Mutation guards also take `PostgreSQL` advisory locks, so they + /// exclude conflicting guards on every replica. Lifecycle locks are + /// process-local. + const fn is_distributed(self) -> bool { + matches!(self, Self::Global | Self::Workspace(_) | Self::Sandbox(..)) + } + } + + /// Holders of each key, changed only while the matching guard is held, so + /// lost exclusion panics instead of passing silently. A counter is `-1` + /// under an exclusive holder and otherwise counts shared holders. A + /// sandbox flag marks the holder of that sandbox key, which does not + /// depend on the workspace. + #[derive(Default)] + struct Occupancy { + global: AtomicIsize, + workspaces: [AtomicIsize; MIX_WORKSPACES], + sandboxes: [AtomicBool; MIX_SANDBOXES], + } + + impl Occupancy { + fn share(counter: &AtomicIsize) { + assert!( + counter.fetch_add(1, Ordering::SeqCst) >= 0, + "shared holder entered under an exclusive holder" + ); + } + + fn exclude(counter: &AtomicIsize) { + assert!( + counter + .compare_exchange(0, -1, Ordering::SeqCst, Ordering::SeqCst) + .is_ok(), + "exclusive holder entered while the key was occupied" + ); + } + + fn claim(sandbox: &AtomicBool) { + assert!( + !sandbox.swap(true, Ordering::SeqCst), + "sandbox entered twice" + ); + } + + fn enter(&self, op: MixOp) { + match op { + MixOp::Global => Self::exclude(&self.global), + MixOp::Workspace(workspace) => { + Self::share(&self.global); + Self::exclude(&self.workspaces[workspace]); + } + MixOp::Sandbox(workspace, sandbox) => { + Self::share(&self.global); + Self::share(&self.workspaces[workspace]); + Self::claim(&self.sandboxes[sandbox]); + } + MixOp::Lifecycle(sandbox) | MixOp::GatedLifecycle(sandbox) => { + Self::share(&self.global); + Self::claim(&self.sandboxes[sandbox]); + } + } + } + + fn leave(&self, op: MixOp) { + match op { + MixOp::Global => self.global.store(0, Ordering::SeqCst), + MixOp::Workspace(workspace) => { + self.workspaces[workspace].store(0, Ordering::SeqCst); + self.global.fetch_sub(1, Ordering::SeqCst); + } + MixOp::Sandbox(workspace, sandbox) => { + self.sandboxes[sandbox].store(false, Ordering::SeqCst); + self.workspaces[workspace].fetch_sub(1, Ordering::SeqCst); + self.global.fetch_sub(1, Ordering::SeqCst); + } + MixOp::Lifecycle(sandbox) | MixOp::GatedLifecycle(sandbox) => { + self.sandboxes[sandbox].store(false, Ordering::SeqCst); + self.global.fetch_sub(1, Ordering::SeqCst); + } + } + } + + fn assert_empty(&self) { + assert_eq!(self.global.load(Ordering::SeqCst), 0); + assert!( + self.workspaces + .iter() + .all(|workspace| workspace.load(Ordering::SeqCst) == 0) + ); + assert!( + self.sandboxes + .iter() + .all(|sandbox| !sandbox.load(Ordering::SeqCst)) + ); + } + } + + /// Enter `op`'s keys, keep them for `hold`, then leave them in reverse + /// order. `local` is the occupancy of the replica that runs `op`; mutation + /// guards also enter `fleet`. The caller holds `op`'s guards throughout. + async fn occupy(local: &Occupancy, fleet: &Occupancy, op: MixOp, hold: Duration) { + local.enter(op); + if op.is_distributed() { + fleet.enter(op); + } + tokio::time::sleep(hold).await; + if op.is_distributed() { + fleet.leave(op); + } + local.leave(op); + } + + /// Run `tasks` tasks of `iterations` random operations each, spread + /// round-robin over `replicas`, holding each operation's locks for 0-2 ms. + /// Every operation must acquire its locks and finish `within`. + async fn run_random_scope_mix( + replicas: &[ComputeRuntime], + prefix: &str, + tasks: usize, + iterations: usize, + within: Duration, + ) { + let names = Arc::new(MixNames::new(prefix)); + let fleet = Arc::new(Occupancy::default()); + let locals: Vec> = replicas.iter().map(|_| Arc::default()).collect(); + let mut rng = StdRng::seed_from_u64(3528); + let mut handles = Vec::new(); + for task in 0..tasks { + let plan: Vec<(MixOp, u64)> = (0..iterations) + .map(|_| { + let op = MixOp::random(&mut rng); + (op, rng.random_range(0..=2)) + }) + .collect(); + let replica = task % replicas.len(); + let runtime = replicas[replica].clone(); + let names = Arc::clone(&names); + let fleet = Arc::clone(&fleet); + let local = Arc::clone(&locals[replica]); + handles.push(tokio::spawn(async move { + for (op, hold_ms) in plan { + let hold = Duration::from_millis(hold_ms); + match op { + MixOp::Global => { + let _guard = runtime + .mutation_guard(MutationScope::Global) + .await + .expect("guard"); + occupy(&local, &fleet, op, hold).await; + } + MixOp::Workspace(workspace) => { + let scope = MutationScope::Workspace(&names.workspaces[workspace]); + let _guard = runtime.mutation_guard(scope).await.expect("guard"); + occupy(&local, &fleet, op, hold).await; + } + MixOp::Sandbox(workspace, sandbox) => { + let scope = MutationScope::sandbox( + &names.workspaces[workspace], + &names.sandboxes[sandbox], + ); + let _guard = runtime.mutation_guard(scope).await.expect("guard"); + occupy(&local, &fleet, op, hold).await; + } + MixOp::Lifecycle(sandbox) => { + let _guard = + runtime.lock_sandbox_local(&names.sandboxes[sandbox]).await; + occupy(&local, &fleet, op, hold).await; + } + MixOp::GatedLifecycle(sandbox) => { + let gate = runtime + .lifecycle_gates + .lock_for(&names.sandboxes[sandbox]) + .await; + let _guard = runtime.lock_sandbox_for_lifecycle(&gate).await; + occupy(&local, &fleet, op, hold).await; + } + } + } + })); + } + + tokio::time::timeout(within, async { + for handle in handles { + handle.await.expect("stress task"); + } + }) + .await + .expect("random scope mix finished without a deadlock"); + for runtime in replicas { + assert_eq!(runtime.mutation_locks.live_entry_count(), 0); + } + fleet.assert_empty(); + for local in &locals { + local.assert_empty(); + } + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn random_scope_mix_never_deadlocks() { + let runtime = test_runtime().await; + run_random_scope_mix(&[runtime], "", 64, 50, Duration::from_secs(20)).await; + } + + #[tokio::test] + async fn unrelated_supervisor_state_update_does_not_wait_for_sandbox_guard() { + let runtime = test_runtime().await; + let mut sandbox = Sandbox { + metadata: Some(openshell_core::proto::datamodel::v1::ObjectMeta { + id: "mutation-guard-unrelated-a".to_string(), + name: "mutation-guard-unrelated-a".to_string(), + workspace: "default".to_string(), + ..Default::default() + }), + ..Default::default() + }; + sandbox.set_phase(SandboxPhase::Provisioning as i32); + runtime.store.put_message(&sandbox).await.unwrap(); + let held = runtime + .mutation_guard(MutationScope::sandbox( + "default", + "mutation-guard-unrelated-b", + )) + .await + .unwrap(); + + tokio::time::timeout( + PROCEEDS_WITHIN, + runtime.supervisor_session_connected("mutation-guard-unrelated-a", "i"), + ) + .await + .expect("an unrelated supervisor update should not wait") + .expect("supervisor session connected"); + let stored = runtime + .store + .get_message::("mutation-guard-unrelated-a") + .await + .unwrap() + .unwrap(); + assert_eq!(stored.phase(), SandboxPhase::Ready as i32); + drop(held); + } + + /// A runtime on its own store connected to `schema`, like one gateway + /// replica. + async fn postgres_runtime(schema: &TestSchema) -> ComputeRuntime { + let store = Arc::new(schema.connect_store().await); + super::super::new_test_runtime_for_driver(store, "test").await + } + + fn stored_sandbox(sandbox_id: &str, workspace: &str) -> Sandbox { + Sandbox { + metadata: Some(openshell_core::proto::datamodel::v1::ObjectMeta { + id: sandbox_id.to_string(), + name: sandbox_id.to_string(), + workspace: workspace.to_string(), + ..Default::default() + }), + ..Default::default() + } + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + #[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] + async fn postgres_compute_guards_random_scope_mix_never_deadlocks() { + let schema = TestSchema::create("mix").await; + let replicas = [ + postgres_runtime(&schema).await, + postgres_runtime(&schema).await, + ]; + // Advisory locks are database-wide, so the mix uses fresh names. + let prefix = format!("{}-", Uuid::new_v4().simple()); + + run_random_scope_mix(&replicas, &prefix, 32, 20, Duration::from_mins(1)).await; + + for replica in &replicas { + replica.store.close().await; + } + schema.drop_schema().await; + } + + #[tokio::test] + #[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] + async fn postgres_compute_guard_by_id_uses_the_sandbox_workspace() { + let schema = TestSchema::create("guard").await; + let replica_a = postgres_runtime(&schema).await; + let replica_b = postgres_runtime(&schema).await; + let workspace = format!("ws-{}", Uuid::new_v4()); + let sandbox_id = format!("sb-{}", Uuid::new_v4()); + replica_a + .store + .put_message(&stored_sandbox(&sandbox_id, &workspace)) + .await + .expect("seed the sandbox"); + + // Only the database locks connect the two replicas. + let held = replica_b + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .expect("workspace guard on replica B"); + let mut by_id = { + let replica_a = replica_a.clone(); + let sandbox_id = sandbox_id.clone(); + tokio::spawn(async move { replica_a.sandbox_mutation_guard_by_id(&sandbox_id).await }) + }; + assert!( + tokio::time::timeout(Duration::from_millis(200), &mut by_id) + .await + .is_err(), + "the by-id guard should wait for the sandbox's workspace" + ); + drop(held); + let guard = tokio::time::timeout(PROCEEDS_WITHIN, by_id) + .await + .expect("the by-id guard should proceed after release") + .expect("guard task") + .expect("by-id guard"); + assert!(guard.is_some(), "the seeded sandbox exists"); + drop(guard); + + assert!( + replica_a + .sandbox_mutation_guard_by_id(&format!("sb-{}", Uuid::new_v4())) + .await + .expect("by-id guard for an unknown sandbox") + .is_none() + ); + + replica_a.store.close().await; + replica_b.store.close().await; + schema.drop_schema().await; + } + + /// Measures lock waits in a paced reconnect burst into one replica: one + /// session every 12 ms (1000 sessions spread over 12 s), all into one + /// receiving replica with the production lock pool. Run it with + /// `--no-capture` to see the wait percentiles. + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + #[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] + async fn postgres_lock_pool_absorbs_a_12ms_reconnect_burst() { + /// Sessions that move to the receiving replica. + const RECONNECTS: usize = 500; + /// One reconnect every 12 ms, as for 1000 sessions spread over 12 s. + const RECONNECT_INTERVAL: Duration = Duration::from_millis(12); + /// Guarded operations per reconnect on the receiver: the pre-ack + /// endpoint-status reset and one endpoint report. + const GUARDED_OPS_PER_RECONNECT: usize = 2; + /// Extra time under the guard, so each critical section takes about + /// 20 ms, like one against a managed database. + const CRITICAL_SECTION_PADDING: Duration = Duration::from_millis(15); + + let schema = TestSchema::create("envelope").await; + // The production lock pool, as on a real receiving replica. + let receiver = postgres_runtime(&schema).await; + let workspace = format!("ws-{}", Uuid::new_v4()); + let mut sandbox_ids = Vec::with_capacity(RECONNECTS); + for _ in 0..RECONNECTS { + let sandbox_id = format!("sb-{}", Uuid::new_v4()); + receiver + .store + .put_message(&stored_sandbox(&sandbox_id, &workspace)) + .await + .expect("seed a sandbox"); + sandbox_ids.push(sandbox_id); + } + + let started = tokio::time::Instant::now(); + let reconnects: Vec<_> = sandbox_ids + .into_iter() + .enumerate() + .map(|(index, sandbox_id)| { + let receiver = receiver.clone(); + let arrives = started + + RECONNECT_INTERVAL * u32::try_from(index).expect("reconnect index fits u32"); + tokio::spawn(async move { + tokio::time::sleep_until(arrives).await; + let mut waits = Vec::with_capacity(GUARDED_OPS_PER_RECONNECT); + for op in 0..GUARDED_OPS_PER_RECONNECT { + let called = tokio::time::Instant::now(); + let guard = receiver + .sandbox_mutation_guard_by_id(&sandbox_id) + .await? + .expect("seeded sandbox"); + waits.push(called.elapsed()); + let sandbox = receiver + .store + .get_message::(&sandbox_id) + .await? + .expect("seeded sandbox"); + receiver + .store + .update_message_cas::( + &sandbox_id, + sandbox.get_resource_version(), + |sandbox| { + sandbox + .metadata + .as_mut() + .expect("sandbox metadata") + .labels + .insert("envelope-op".to_string(), op.to_string()); + }, + ) + .await?; + tokio::time::sleep(CRITICAL_SECTION_PADDING).await; + drop(guard); + } + Ok::<_, PersistenceError>(waits) + }) + }) + .collect(); + let mut waits = Vec::with_capacity(RECONNECTS * GUARDED_OPS_PER_RECONNECT); + for reconnect in reconnects { + match reconnect.await.expect("reconnect task") { + Ok(reconnect_waits) => waits.extend(reconnect_waits), + Err(error) => panic!("a guarded reconnect operation failed: {error:?}"), + } + } + + waits.sort_unstable(); + let percentile = |percent: usize| waits[(waits.len() * percent).div_ceil(100) - 1]; + let (p50, p99) = (percentile(50), percentile(99)); + let max = waits[waits.len() - 1]; + eprintln!( + "reconnect burst: {} guarded ops from {RECONNECTS} reconnects {RECONNECT_INTERVAL:?} \ + apart into one receiver: lock wait p50 {p50:?}, p99 {p99:?}, max {max:?}", + waits.len() + ); + // Waits must stay far from the lock timeout, where requests fail. + assert!( + p99 * 5 < MUTATION_LOCK_TIMEOUT, + "p99 lock wait {p99:?} is too close to the {MUTATION_LOCK_TIMEOUT:?} timeout" + ); + + receiver.store.close().await; + schema.drop_schema().await; + } +} diff --git a/crates/openshell-server/src/compute/provisioning_deadline.rs b/crates/openshell-server/src/compute/provisioning_deadline.rs index 1724e3c984..7b7608d38d 100644 --- a/crates/openshell-server/src/compute/provisioning_deadline.rs +++ b/crates/openshell-server/src/compute/provisioning_deadline.rs @@ -447,7 +447,7 @@ impl super::ComputeRuntime { pub(super) async fn reconcile_provisioning_deadlines(&self, now_ms: i64) -> Result<(), String> { use crate::persistence::{ObjectListQuery, ObjectType}; use openshell_core::{ - ObjectId, + ObjectId, ObjectWorkspace, proto::{Sandbox, SandboxPhase}, }; use prost::Message; @@ -469,7 +469,9 @@ impl super::ComputeRuntime { // Expiration can fence Starting while its driver RPC owns the // lifecycle gate. Cleanup waits for that gate; Error never waits // for compute I/O, matching the existing driver-observation fence. - let global = self.sync_lock.clone().lock_owned().await; + let local_guard = self + .lock_sandbox_local_in_workspace(candidate.object_workspace(), &record.id) + .await; let Some(mut current) = self .store .get_message::(&record.id) @@ -495,7 +497,7 @@ impl super::ComputeRuntime { if let Some(expired) = self.claim_provisioning_timeout(¤t, now_ms).await? { current = expired; } - drop(global); + drop(local_guard); if timed_out(¤t) && current .status @@ -510,10 +512,9 @@ impl super::ComputeRuntime { .is_none_or(|t| t <= now_ms) }) { - let Ok(guard) = self.lifecycle_gates.gate_for(&record.id).try_lock_owned() else { + let Some(gate) = self.lifecycle_gates.try_lock_for(&record.id) else { continue; }; - let gate = super::SandboxLifecycleGuard { _guard: guard }; let runtime = self.clone(); tokio::spawn(async move { if let Err(error) = runtime.reclaim_provisioning_timeout(¤t, &gate).await @@ -527,7 +528,7 @@ impl super::ComputeRuntime { } /// Claim expiration durably before touching the backend. The caller owns the - /// global configuration guard; CAS fences concurrent lifecycle operations. + /// sandbox's local mutation lock; CAS fences concurrent lifecycle operations. /// The separate cleanup step also requires the per-sandbox lifecycle gate. pub(crate) async fn claim_provisioning_timeout( &self, @@ -621,7 +622,7 @@ impl super::ComputeRuntime { use openshell_core::proto::compute::v1::StopSandboxRequest; use openshell_core::{ObjectId, ObjectName}; let current = { - let _global_guard = self.lock_global_for_lifecycle(lifecycle_guard).await; + let _sandbox_guard = self.lock_sandbox_for_lifecycle(lifecycle_guard).await; self.store .get_message::(expired.object_id()) .await @@ -674,7 +675,7 @@ impl super::ComputeRuntime { // Cross-replica cleanup claim. A replacement leader waits longer than // the bounded driver call before retrying an interrupted reclamation. { - let _global_guard = self.lock_global_for_lifecycle(lifecycle_guard).await; + let _sandbox_guard = self.lock_sandbox_for_lifecycle(lifecycle_guard).await; self.store .update_message_cas::( expired.object_id(), @@ -716,7 +717,7 @@ impl super::ComputeRuntime { let reclaimed = !pending_before_stop && (matches!(&result, Ok(Ok(_))) || matches!(&result, Ok(Err(error)) if error.code() == tonic::Code::NotFound)); - let _global_guard = self.lock_global_for_lifecycle(lifecycle_guard).await; + let _sandbox_guard = self.lock_sandbox_for_lifecycle(lifecycle_guard).await; let Some(current) = self .store .get_message::(&sandbox_id) diff --git a/crates/openshell-server/src/gateway_metrics.rs b/crates/openshell-server/src/gateway_metrics.rs index 28adda1712..0cd4ffbac8 100644 --- a/crates/openshell-server/src/gateway_metrics.rs +++ b/crates/openshell-server/src/gateway_metrics.rs @@ -29,9 +29,11 @@ pub const RELAY_PENDING_CAPACITY: &str = "openshell_server_relay_pending_capacit pub const RELAY_REJECTED_TOTAL: &str = "openshell_server_relay_rejected_total"; pub const RELAY_EXPIRED_TOTAL: &str = "openshell_server_relay_expired_total"; pub const ROUTED_REQUEST_ATTEMPTS_TOTAL: &str = "openshell_server_routed_request_attempts_total"; +pub const MUTATION_LOCK_TIMEOUTS_TOTAL: &str = "openshell_server_mutation_lock_timeouts_total"; // Histograms (explicit buckets, see BUCKETED_HISTOGRAMS) pub const RELAY_CLAIM_DURATION_SECONDS: &str = "openshell_server_relay_claim_duration_seconds"; pub const PEER_REQUEST_DURATION_SECONDS: &str = "openshell_server_peer_request_duration_seconds"; +pub const MUTATION_LOCK_WAIT_SECONDS: &str = "openshell_server_mutation_lock_wait_seconds"; const LABEL_REASON: &str = "reason"; const LABEL_OPERATION: &str = "operation"; @@ -39,17 +41,21 @@ const LABEL_OUTCOME: &str = "outcome"; const LABEL_GRPC_CODE: &str = "grpc_code"; const LABEL_RELAY_KIND: &str = "relay_kind"; const LABEL_ROUTE: &str = "route"; +const LABEL_SCOPE: &str = "scope"; /// Buckets for the new latency histograms, 1 ms to 15 s. The top buckets cover the 10 s relay -/// claim timeout and the 15 s routed-relay wait. +/// claim and lock timeouts and the 15 s routed-relay wait. const LATENCY_BUCKETS_SECONDS: [f64; 14] = [ 0.001, 0.0025, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 15.0, ]; /// Only these names render as Prometheus histograms. Every existing `*_duration_seconds` metric /// keeps its summary format, so current dashboards are unaffected. -const BUCKETED_HISTOGRAMS: [&str; 2] = - [RELAY_CLAIM_DURATION_SECONDS, PEER_REQUEST_DURATION_SECONDS]; +const BUCKETED_HISTOGRAMS: [&str; 3] = [ + RELAY_CLAIM_DURATION_SECONDS, + PEER_REQUEST_DURATION_SECONDS, + MUTATION_LOCK_WAIT_SECONDS, +]; /// Protocol the supervisor is asked to relay. Never label metrics with the target address. #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -133,6 +139,26 @@ impl PeerRpc { } } +/// Mutation lock scope kind. The platform scope ("" workspace) maps to `Global`. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum LockScope { + Global, + Workspace, + Sandbox, +} + +impl LockScope { + pub const ALL: [Self; 3] = [Self::Global, Self::Workspace, Self::Sandbox]; + + pub const fn label(self) -> &'static str { + match self { + Self::Global => "global", + Self::Workspace => "workspace", + Self::Sandbox => "sandbox", + } + } +} + /// Where a routed attempt ended. A relay succeeds only when the supervisor claims it, on either /// route, so the values mean the same thing for local and peer attempts. #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -240,6 +266,11 @@ pub fn describe_and_initialize(relay: RelayCapacity) { Unit::Count, "Pending relay channels dropped because the supervisor did not connect back in time." ); + describe_counter!( + MUTATION_LOCK_TIMEOUTS_TOTAL, + Unit::Count, + "Mutation lock acquisitions that timed out." + ); describe_histogram!( RELAY_CLAIM_DURATION_SECONDS, Unit::Seconds, @@ -250,6 +281,11 @@ pub fn describe_and_initialize(relay: RelayCapacity) { Unit::Seconds, "Latency of outbound requests to the owning replica. For relays, until the owner's supervisor claimed the relay." ); + describe_histogram!( + MUTATION_LOCK_WAIT_SECONDS, + Unit::Seconds, + "Time spent acquiring the mutation lock for a scope." + ); // `increment(0)` registers a series without overwriting a value recorded earlier. gauge!(SUPERVISOR_SESSIONS).increment(0.0); @@ -272,6 +308,9 @@ pub fn describe_and_initialize(relay: RelayCapacity) { counter!(RELAY_REJECTED_TOTAL, LABEL_REASON => reason.label()).increment(0); } counter!(RELAY_EXPIRED_TOTAL).increment(0); + for scope in LockScope::ALL { + counter!(MUTATION_LOCK_TIMEOUTS_TOTAL, LABEL_SCOPE => scope.label()).increment(0); + } for rpc in PeerRpc::ALL { if rpc == PeerRpc::Relay { continue; @@ -339,6 +378,20 @@ pub fn record_relay_claimed(waited: Duration) { histogram!(RELAY_CLAIM_DURATION_SECONDS).record(waited); } +/// Time to acquire every key of one mutation guard (local registry plus Postgres), recorded +/// on success only. +pub fn record_lock_wait(scope: LockScope, waited: Duration) { + histogram!(MUTATION_LOCK_WAIT_SECONDS, LABEL_SCOPE => scope.label()).record(waited); +} + +/// A guard acquisition that timed out: a local wait, a full lock pool, too little time left to +/// open a lock connection, or Postgres `lock_timeout` (SQLSTATE 55P03). A lock connection that +/// Postgres does not open with at least `LOCK_CONNECTION_MIN_BUDGET` left is not counted. RPC +/// callers return the timeout as `Status::unavailable`. +pub fn record_lock_timeout(scope: LockScope) { + counter!(MUTATION_LOCK_TIMEOUTS_TOTAL, LABEL_SCOPE => scope.label()).increment(1); +} + /// Counts one local relay setup or outbound peer attempt exactly once, and times peer requests. /// A relay succeeds when the supervisor claims it, on either route. Dropping an unfinished /// timer (the caller gave up) records `local_error` / `cancelled`. @@ -503,6 +556,18 @@ mod tests { 0, ), ("openshell_server_relay_expired_total", 0), + ( + "openshell_server_mutation_lock_timeouts_total{scope=\"global\"}", + 0, + ), + ( + "openshell_server_mutation_lock_timeouts_total{scope=\"workspace\"}", + 0, + ), + ( + "openshell_server_mutation_lock_timeouts_total{scope=\"sandbox\"}", + 0, + ), ] { assert_eq!(metrics.value(series), Some(expected), "{series}"); } @@ -579,6 +644,7 @@ mod tests { LABEL_OUTCOME => "success" ) .record(sample); + histogram!(MUTATION_LOCK_WAIT_SECONDS, LABEL_SCOPE => "sandbox").record(sample); histogram!( "openshell_server_grpc_request_duration_seconds", "method" => "ListSandboxes", diff --git a/crates/openshell-server/src/grpc/mod.rs b/crates/openshell-server/src/grpc/mod.rs index 5ac58307a4..156894804a 100644 --- a/crates/openshell-server/src/grpc/mod.rs +++ b/crates/openshell-server/src/grpc/mod.rs @@ -78,8 +78,10 @@ use crate::ServerState; /// Map a `PersistenceError` to an appropriate gRPC `Status`. /// /// CAS conflicts (optimistic concurrency failures) are mapped to `ABORTED` -/// to signal that the client should retry with fresh data. Other persistence -/// errors are mapped to `INTERNAL`. +/// to signal that the client should retry with fresh data. Mutation lock +/// timeouts are mapped to `UNAVAILABLE` with a retry delay: the mutation lock +/// was not acquired, so this request's guarded writes did not run. Other +/// persistence errors are mapped to `INTERNAL`. pub fn persistence_error_to_status( err: crate::persistence::PersistenceError, operation: &str, @@ -97,6 +99,11 @@ pub fn persistence_error_to_status( ), current_resource_version, ), + PersistenceError::LockTimeout(_) => openshell_core::rpc_error::unavailable( + "MUTATION_LOCK_TIMEOUT", + format!("{operation} timed out waiting for a concurrent mutation; retry the request"), + std::time::Duration::from_secs(1), + ), other => Status::internal(format!("{operation} failed: {other}")), } } @@ -183,6 +190,10 @@ struct StoredSettings { /// loaded from `ObjectRecord` and used for optimistic concurrency control. #[serde(skip)] resource_version: u64, + /// Database id of the loaded row. Not persisted; a save aborts when the + /// row was deleted and recreated under the same name since the load. + #[serde(skip)] + record_id: Option, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -1159,6 +1170,32 @@ mod tests { assert!(gpu.count_selection_supported); } + #[test] + fn persistence_error_to_status_maps_mutation_lock_timeout_to_unavailable() { + let status = persistence_error_to_status( + crate::persistence::PersistenceError::LockTimeout( + "waiting for a PostgreSQL advisory lock".into(), + ), + "acquire provider mutation lock", + ); + + assert_eq!(status.code(), tonic::Code::Unavailable); + assert_eq!( + status.message(), + "acquire provider mutation lock timed out waiting for a concurrent mutation; \ + retry the request" + ); + let details = openshell_core::rpc_error::decode_details(&status).expect("error details"); + assert_eq!( + details.error_info().expect("error info").reason, + "MUTATION_LOCK_TIMEOUT" + ); + assert_eq!( + details.retry_info().expect("retry info").retry_delay, + Some(std::time::Duration::from_secs(1)) + ); + } + #[test] fn public_resource_capabilities_preserves_absence() { let absent: Option = None; diff --git a/crates/openshell-server/src/grpc/policy.rs b/crates/openshell-server/src/grpc/policy.rs index bf115370b4..70067fb01b 100644 --- a/crates/openshell-server/src/grpc/policy.rs +++ b/crates/openshell-server/src/grpc/policy.rs @@ -25,6 +25,7 @@ pub use endpoint_status::{ use crate::ServerState; use crate::auth::principal::Principal; use crate::auth::workspace_authz::{MinWorkspaceRole, require_platform_admin}; +use crate::compute::MutationScope; use crate::pagination::Pagination; use crate::persistence::{ DraftChunkRecord, ObjectId, ObjectListQuery, ObjectName, ObjectType, ObjectWorkspace, @@ -1234,7 +1235,7 @@ fn background_pending_refreshes() -> &'static std::sync::Mutex sandbox lock order for all global policy mutations. + // The global guard taken at the top of this branch excludes every + // sandbox-scoped settings, policy, and report mutation. let mut global_settings = load_global_settings(state.store.as_ref()).await?; let provider_composition_was_enabled = provider_policy_composition_enabled_in(&global_settings)?; @@ -3943,10 +3948,13 @@ async fn handle_update_config_inner( let mut response_annotations = sandbox_metadata_annotations(&sandbox); if has_setting { - let _settings_guard = state.settings_mutex.lock().await; - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::persistence_error_to_status(error, "acquire policy mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::sandbox(&workspace, &sandbox_id)) + .await + .map_err(|error| { + super::persistence_error_to_status(error, "acquire policy mutation lock") + })?; if key == POLICY_SETTING_KEY { return Err(Status::invalid_argument( @@ -3967,6 +3975,7 @@ async fn handle_update_config_inner( let mut sandbox_settings = load_sandbox_settings(state.store.as_ref(), &workspace, sandbox.object_name()) .await?; + ensure_sandbox_keeps_name(state, &sandbox).await?; let removed = sandbox_settings.settings.remove(key).is_some(); if removed { sandbox_settings.revision = sandbox_settings.revision.wrapping_add(1); @@ -4011,6 +4020,7 @@ async fn handle_update_config_inner( let mut sandbox_settings = load_sandbox_settings(state.store.as_ref(), &workspace, sandbox.object_name()).await?; + ensure_sandbox_keeps_name(state, &sandbox).await?; let changed = upsert_setting_value(&mut sandbox_settings.settings, key, stored); if changed { sandbox_settings.revision = sandbox_settings.revision.wrapping_add(1); @@ -4041,9 +4051,13 @@ async fn handle_update_config_inner( )); } - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::persistence_error_to_status(error, "acquire policy mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::sandbox(&workspace, &sandbox_id)) + .await + .map_err(|error| { + super::persistence_error_to_status(error, "acquire policy mutation lock") + })?; if has_merge_ops { let global_settings = load_global_settings(state.store.as_ref()).await?; if global_settings.settings.contains_key(POLICY_SETTING_KEY) { @@ -4580,9 +4594,16 @@ pub(super) async fn handle_report_sandbox_configuration( if reported == ConfigurationAdmissionState::Unspecified { return Err(Status::invalid_argument("admission state is required")); } - let _guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::persistence_error_to_status(error, "acquire configuration admission lock") - })?; + let Some(_mutation_guard) = state + .compute + .sandbox_mutation_guard_by_id(&sandbox_id) + .await + .map_err(|error| { + super::persistence_error_to_status(error, "acquire configuration admission lock") + })? + else { + return Err(Status::not_found("sandbox not found")); + }; let mut sandbox = state .store .get_message::(&sandbox_id) @@ -4782,9 +4803,16 @@ pub(super) async fn handle_report_policy_status( .supersede_older_policies(&req.sandbox_id, version) .await; - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::persistence_error_to_status(error, "acquire policy mutation lock") - })?; + let Some(_mutation_guard) = state + .compute + .sandbox_mutation_guard_by_id(&req.sandbox_id) + .await + .map_err(|error| { + super::persistence_error_to_status(error, "acquire policy mutation lock") + })? + else { + return Err(Status::not_found("sandbox not found")); + }; let sandbox = state .store .get_message::(&req.sandbox_id) @@ -7443,6 +7471,30 @@ pub(super) async fn save_global_settings( .await } +/// Sandbox settings are keyed by name, and a sandbox's mutation guard does not +/// exclude a new sandbox that takes the name once this one is deleted. Called +/// after loading the settings: sandbox IDs are never reused, so if the sandbox +/// still exists under its name, the loaded settings are its own. +async fn ensure_sandbox_keeps_name(state: &ServerState, sandbox: &Sandbox) -> Result<(), Status> { + let current = state + .store + .get_message::(sandbox.object_id()) + .await + .map_err(|e| Status::internal(format!("fetch sandbox failed: {e}")))?; + match current { + Some(current) + if current.object_name() == sandbox.object_name() + && current.object_workspace() == sandbox.object_workspace() => + { + Ok(()) + } + _ => Err(Status::not_found(format!( + "sandbox '{}' was deleted", + sandbox.object_name() + ))), + } +} + pub(super) async fn load_sandbox_settings( store: &Store, workspace: &str, @@ -7481,6 +7533,7 @@ async fn load_settings_record( let mut settings = serde_json::from_slice::(&record.payload) .map_err(|e| Status::internal(format!("decode settings payload failed: {e}")))?; settings.resource_version = record.resource_version; + settings.record_id = Some(record.id.clone()); for key in settings.settings.keys() { settings .change_clocks @@ -7531,6 +7584,17 @@ async fn save_settings_record( .await .map_err(|e| Status::internal(format!("fetch settings for CAS failed: {e}")))? .ok_or_else(|| Status::not_found("settings disappeared since load"))?; + // Settings are keyed by name. A row recreated under the same name, for + // a sandbox that reused it, restarts at the loaded resource version. + if settings + .record_id + .as_deref() + .is_some_and(|loaded| loaded != existing.id) + { + return Err(Status::aborted( + "settings were replaced concurrently; please retry", + )); + } ( existing.id, @@ -20855,6 +20919,52 @@ mod tests { ); } + #[tokio::test] + async fn sandbox_setting_update_does_not_wait_for_unrelated_sandbox_guard() { + let state = test_server_state().await; + state + .store + .put_message(&test_sandbox( + "sb-setting-target", + "setting-target", + ProtoSandboxPolicy::default(), + Vec::new(), + )) + .await + .unwrap(); + let unrelated_guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", "sb-setting-unrelated")) + .await + .unwrap(); + + tokio::time::timeout( + std::time::Duration::from_secs(5), + handle_update_config( + &state, + authed_request(UpdateConfigRequest { + sandbox: "setting-target".to_string(), + workspace_scope: Some(openshell_core::proto::workspace_selector( + "default".to_string(), + )), + setting_key: "ocsf_json_enabled".to_string(), + setting_value: Some(SettingValue { + value: Some(setting_value::Value::BoolValue(true)), + }), + ..Default::default() + }), + ), + ) + .await + .expect("a sandbox setting update should not wait for an unrelated sandbox") + .expect("sandbox setting update succeeds"); + let settings = load_sandbox_settings(state.store.as_ref(), "default", "setting-target") + .await + .unwrap(); + assert!(settings.settings.contains_key("ocsf_json_enabled")); + drop(unrelated_guard); + } + #[tokio::test] async fn update_config_global_policy_rejects_reserved_provider_key() { let state = test_server_state().await; @@ -21605,6 +21715,87 @@ mod tests { ); } + #[tokio::test] + async fn sandbox_settings_save_aborts_when_row_was_recreated() { + let store = test_store().await; + let sandbox_name = "reused-name"; + let mut original = StoredSettings::default(); + original + .settings + .insert("a_key".to_string(), StoredSettingValue::Bool(true)); + save_sandbox_settings(&store, "default", sandbox_name, &original) + .await + .unwrap(); + let mut loaded = load_sandbox_settings(&store, "default", sandbox_name) + .await + .unwrap(); + + // A new sandbox reuses the name: its settings row restarts at the + // version the stale load saw. + store + .delete_by_name(SANDBOX_SETTINGS_OBJECT_TYPE, "default", sandbox_name) + .await + .unwrap(); + let mut replacement = StoredSettings::default(); + replacement + .settings + .insert("c_key".to_string(), StoredSettingValue::Bool(true)); + save_sandbox_settings(&store, "default", sandbox_name, &replacement) + .await + .unwrap(); + let recreated = load_sandbox_settings(&store, "default", sandbox_name) + .await + .unwrap(); + assert_eq!(recreated.resource_version, loaded.resource_version); + + loaded + .settings + .insert("r_key".to_string(), StoredSettingValue::Bool(true)); + let error = save_sandbox_settings(&store, "default", sandbox_name, &loaded) + .await + .unwrap_err(); + assert_eq!(error.code(), Code::Aborted); + let current = load_sandbox_settings(&store, "default", sandbox_name) + .await + .unwrap(); + assert!(current.settings.contains_key("c_key")); + assert!(!current.settings.contains_key("r_key")); + } + + #[tokio::test] + async fn sandbox_settings_owner_check_rejects_a_reused_name() { + let state = test_server_state().await; + let original = test_sandbox( + "sb-original", + "reused-name", + ProtoSandboxPolicy::default(), + Vec::new(), + ); + state.store.put_message(&original).await.unwrap(); + ensure_sandbox_keeps_name(&state, &original).await.unwrap(); + + state + .store + .delete(Sandbox::object_type(), original.object_id()) + .await + .unwrap(); + let replacement = test_sandbox( + "sb-replacement", + "reused-name", + ProtoSandboxPolicy::default(), + Vec::new(), + ); + state.store.put_message(&replacement).await.unwrap(); + + let error = ensure_sandbox_keeps_name(&state, &original) + .await + .unwrap_err(); + assert_eq!(error.code(), Code::NotFound); + ensure_sandbox_keeps_name(&state, &replacement) + .await + .unwrap(); + } + #[tokio::test] async fn concurrent_global_setting_mutations_are_serialized() { let store = Arc::new(test_store().await); diff --git a/crates/openshell-server/src/grpc/policy/endpoint_status.rs b/crates/openshell-server/src/grpc/policy/endpoint_status.rs index 462099198b..7576251b19 100644 --- a/crates/openshell-server/src/grpc/policy/endpoint_status.rs +++ b/crates/openshell-server/src/grpc/policy/endpoint_status.rs @@ -10,11 +10,15 @@ use super::{ deterministic_policy_hash, load_global_settings, policy_static_credential_endpoint_bindings, }; use crate::ServerState; -use crate::persistence::{ObjectId, ObjectWorkspace}; +use crate::compute::MutationScope; +use crate::persistence::{ + ObjectCursor, ObjectId, ObjectListQuery, ObjectWorkspace, PersistenceError, +}; use crate::policy_store::PolicyStoreExt; use crate::provider_profile_sources::EffectiveProviderProfileCatalog; use crate::supervisor_owner::{OWNER_TTL, SupervisorOwnerIndex}; use crate::supervisor_session::EndpointReportCursor; +use futures::TryStreamExt; use openshell_core::GetResourceVersion; use openshell_core::endpoint_status::initial_endpoint_status; use openshell_core::mcp::is_mcp_protocol; @@ -30,6 +34,11 @@ use tonic::{Request, Response, Status}; use tracing::warn; const ENDPOINT_STARTUP_RECONCILIATION_PAGE_SIZE: u32 = 100; +/// Sandboxes reconciled at once, matching the four-connection mutation lock +/// pool. +const ENDPOINT_STARTUP_RECONCILIATION_CONCURRENCY: usize = 4; +/// Guarded attempts per sandbox before a concurrent write fails startup. +const ENDPOINT_STARTUP_RECONCILIATION_ATTEMPTS: usize = 5; const ENDPOINT_DISCONNECT_RETRY_INITIAL_BACKOFF: std::time::Duration = std::time::Duration::from_millis(100); const ENDPOINT_DISCONNECT_RETRY_MAX_BACKOFF: std::time::Duration = @@ -111,9 +120,19 @@ async fn handle_report_endpoint_status_inner( // Session validation, configuration derivation, and persistence share the // sandbox mutation boundary. A newly registered supervisor can therefore // invalidate its predecessor before any stale report reaches the CAS. - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::super::persistence_error_to_status(error, "acquire endpoint status mutation lock") - })?; + let Some(_mutation_guard) = state + .compute + .sandbox_mutation_guard_by_id(&req.sandbox_id) + .await + .map_err(|error| { + super::super::persistence_error_to_status( + error, + "acquire endpoint status mutation lock", + ) + })? + else { + return Err(Status::not_found("sandbox not found")); + }; if !state .supervisor_sessions .is_endpoint_status_authority(&req.sandbox_id, &req.supervisor_session_id) @@ -362,9 +381,19 @@ pub async fn reset_endpoint_status_for_supervisor_session( sandbox_id: &str, supervisor_session_id: &str, ) -> Result<(), Status> { - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::super::persistence_error_to_status(error, "acquire endpoint status mutation lock") - })?; + let Some(_mutation_guard) = state + .compute + .sandbox_mutation_guard_by_id(sandbox_id) + .await + .map_err(|error| { + super::super::persistence_error_to_status( + error, + "acquire endpoint status mutation lock", + ) + })? + else { + return Err(Status::not_found("sandbox not found")); + }; if !state .supervisor_sessions .is_current_session(sandbox_id, supervisor_session_id) @@ -404,22 +433,33 @@ pub async fn reset_endpoint_status_for_supervisor_session( /// Reset endpoint observations after the active supervisor stream disconnects. /// /// A concurrently registered replacement owns its own pre-acknowledgement -/// reset, so this path leaves that session's cursor alone. +/// reset, so this path leaves that session's cursor alone. A replacement that +/// already exists is detected before the mutation guard, so the reset does not +/// wait behind other mutations of the sandbox. pub async fn reset_endpoint_status_after_supervisor_disconnect( state: &Arc, sandbox_id: &str, ) -> Result<(), Status> { - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::super::persistence_error_to_status(error, "acquire endpoint status mutation lock") - })?; - if state - .supervisor_sessions - .current_session_id(sandbox_id) - .is_some() - { + if disconnect_reset_is_superseded(state, sandbox_id).await? { return Ok(()); } - if has_fresh_shared_owner(state, sandbox_id).await? { + let Some(_mutation_guard) = state + .compute + .sandbox_mutation_guard_by_id(sandbox_id) + .await + .map_err(|error| { + super::super::persistence_error_to_status( + error, + "acquire endpoint status mutation lock", + ) + })? + else { + return Err(Status::not_found("sandbox not found")); + }; + // Re-check under the guard: a replacement's pre-acknowledgement reset runs + // under the same sandbox key, and resetting after it would wipe the + // replacement's fresh evidence. + if disconnect_reset_is_superseded(state, sandbox_id).await? { return Ok(()); } let sandbox = state @@ -484,68 +524,153 @@ pub async fn retry_endpoint_status_after_supervisor_disconnect( /// Supervisor sessions are intentionally process-local. This reconciliation /// runs before gateway listeners are bound. A fresh shared owner preserves its /// evidence; records without one are reset so stale success is never served. +/// +/// Each sandbox is reset under its own mutation guard, so other replicas keep +/// mutating unrelated sandboxes during the scan. Keyset paging keeps the scan +/// stable when sandboxes are deleted mid-scan. pub async fn invalidate_endpoint_status_on_startup(state: &Arc) -> Result<(), Status> { - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::super::persistence_error_to_status( - error, - "acquire endpoint status startup reconciliation lock", - ) - })?; - let mut offset = 0; + let mut cursor: Option = None; loop { - let sandboxes = state + let page = state .store - .list_all_messages::(ENDPOINT_STARTUP_RECONCILIATION_PAGE_SIZE, offset) + .list_message_page::( + ObjectListQuery::AllWorkspaces, + cursor.as_ref(), + ENDPOINT_STARTUP_RECONCILIATION_PAGE_SIZE, + ) .await .map_err(|error| { Status::internal(format!( "list sandboxes for tool server endpoint-status startup reconciliation failed: {error}" )) })?; - if sandboxes.is_empty() { + // Each future holds at most its own sandbox guard and never waits on + // another, so running them in one task cannot deadlock. + futures::stream::iter( + page.messages + .iter() + .filter(|sandbox| has_endpoint_status(sandbox)) + .map(Ok::<_, Status>), + ) + .try_for_each_concurrent(ENDPOINT_STARTUP_RECONCILIATION_CONCURRENCY, |candidate| { + invalidate_sandbox_endpoint_status_on_startup(state, candidate) + }) + .await?; + let Some(next_cursor) = page.next_cursor else { return Ok(()); - } + }; + cursor = Some(next_cursor); + } +} - for sandbox in &sandboxes { - let has_endpoint_status = sandbox - .status - .as_ref() - .is_some_and(|status| !status.endpoint_statuses.is_empty()); - if !has_endpoint_status { - continue; - } - let sandbox_id = sandbox.object_id(); - if has_fresh_shared_owner(state, sandbox_id).await? { - continue; +async fn invalidate_sandbox_endpoint_status_on_startup( + state: &Arc, + candidate: &Sandbox, +) -> Result<(), Status> { + // A live owner keeps its evidence, so it costs no guard. + if has_fresh_shared_owner(state, candidate.object_id()).await? { + return Ok(()); + } + retry_startup_reconciliation(|| invalidate_sandbox_endpoint_status_once(state, candidate)).await +} + +/// Outcome of one guarded startup-reconciliation attempt. +enum StartupAttempt { + Done, + /// The write hit a concurrent change, a row deleted after the re-read, or + /// another database error; re-read and try again. + Retry(PersistenceError), +} + +/// Run `attempt` until it is done, at most +/// `ENDPOINT_STARTUP_RECONCILIATION_ATTEMPTS` times. An error from `attempt` +/// is returned without a retry. +async fn retry_startup_reconciliation(mut attempt: F) -> Result<(), Status> +where + F: FnMut() -> Fut, + Fut: Future>, +{ + let mut attempts = 1; + loop { + match attempt().await? { + StartupAttempt::Done => return Ok(()), + StartupAttempt::Retry(error) + if attempts >= ENDPOINT_STARTUP_RECONCILIATION_ATTEMPTS => + { + return Err(super::super::persistence_error_to_status( + error, + "invalidate tool server endpoint status during gateway startup", + )); } - let expected_resource_version = sandbox.get_resource_version(); - let updated = state - .store - .update_message_cas::( - sandbox_id, - expected_resource_version, - invalidate_endpoint_status_without_session, - ) - .await - .map_err(|error| { - super::super::persistence_error_to_status( - error, - "invalidate tool server endpoint status during gateway startup", - ) - })?; - state.sandbox_index.update_from_sandbox(&updated); + StartupAttempt::Retry(_) => attempts += 1, } + } +} - let page_len = sandboxes.len() as u32; - if page_len < ENDPOINT_STARTUP_RECONCILIATION_PAGE_SIZE { - return Ok(()); - } - offset = offset.checked_add(page_len).ok_or_else(|| { - Status::internal( - "sandbox pagination overflow during tool server endpoint-status reconciliation", +async fn invalidate_sandbox_endpoint_status_once( + state: &Arc, + candidate: &Sandbox, +) -> Result { + let sandbox_id = candidate.object_id(); + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::sandbox( + candidate.object_workspace(), + sandbox_id, + )) + .await + .map_err(|error| { + super::super::persistence_error_to_status( + error, + "acquire endpoint status startup reconciliation lock", ) })?; + // A supervisor may have connected to another replica while this waited. + if has_fresh_shared_owner(state, sandbox_id).await? { + return Ok(StartupAttempt::Done); } + let Some(current) = state + .store + .get_message::(sandbox_id) + .await + .map_err(|error| Status::internal(format!("fetch sandbox failed: {error}")))? + else { + return Ok(StartupAttempt::Done); + }; + if !has_endpoint_status(¤t) { + return Ok(StartupAttempt::Done); + } + match state + .store + .update_message_cas::( + sandbox_id, + current.get_resource_version(), + invalidate_endpoint_status_without_session, + ) + .await + { + Ok(updated) => { + state.sandbox_index.update_from_sandbox(&updated); + Ok(StartupAttempt::Done) + } + // Lifecycle writers on other replicas take no distributed guard, and a + // delete between the re-read and the write surfaces as a database + // error. Retry with a fresh read after this guard drops. + Err(error @ (PersistenceError::Conflict { .. } | PersistenceError::Database(_))) => { + Ok(StartupAttempt::Retry(error)) + } + Err(error) => Err(super::super::persistence_error_to_status( + error, + "invalidate tool server endpoint status during gateway startup", + )), + } +} + +fn has_endpoint_status(sandbox: &Sandbox) -> bool { + sandbox + .status + .as_ref() + .is_some_and(|status| !status.endpoint_statuses.is_empty()) } async fn has_fresh_shared_owner( @@ -559,6 +684,22 @@ async fn has_fresh_shared_owner( .map_err(|error| Status::unavailable(format!("resolve supervisor owner failed: {error}"))) } +/// True when a replacement supervisor session (local, or a fresh owner on a +/// peer) now owns endpoint observation, so a disconnect must not reset. +async fn disconnect_reset_is_superseded( + state: &Arc, + sandbox_id: &str, +) -> Result { + if state + .supervisor_sessions + .current_session_id(sandbox_id) + .is_some() + { + return Ok(true); + } + has_fresh_shared_owner(state, sandbox_id).await +} + fn invalidate_endpoint_status_without_session(sandbox: &mut Sandbox) { let Some(status) = sandbox.status.as_mut() else { return; diff --git a/crates/openshell-server/src/grpc/policy/endpoint_status_tests.rs b/crates/openshell-server/src/grpc/policy/endpoint_status_tests.rs index 4c951863a0..c974683552 100644 --- a/crates/openshell-server/src/grpc/policy/endpoint_status_tests.rs +++ b/crates/openshell-server/src/grpc/policy/endpoint_status_tests.rs @@ -4,13 +4,17 @@ use super::super::tests::{mcp_policy_with_versions, test_sandbox, with_sandbox}; use super::super::{handle_get_sandbox_config, handle_report_policy_status, handle_update_config}; use super::*; +use crate::grpc::OpenShellService; use crate::grpc::test_support::{authed_request, test_server_state}; use openshell_core::endpoint_status::endpoint_id; +use openshell_core::proto::open_shell_client::OpenShellClient; +use openshell_core::proto::open_shell_server::OpenShellServer; use openshell_core::proto::{ EndpointObservation, GetSandboxConfigRequest, GetSandboxRequest, NetworkEndpoint, NetworkPolicyRule, PolicyStatus, ReportPolicyStatusRequest, SandboxCondition, SandboxPhase, - UpdateConfigRequest, + SupervisorHello, SupervisorMessage, UpdateConfigRequest, supervisor_message, }; +use tokio_stream::wrappers::TcpListenerStream; use tonic::Code; fn timestamp(value: &str) -> prost_types::Timestamp { @@ -408,7 +412,11 @@ async fn global_policy_update_waits_for_endpoint_report_guard() { let before = load_global_settings(state.store.as_ref()) .await .expect("read settings before update"); - let guard = state.compute.sandbox_sync_guard().await.unwrap(); + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", "endpoint-report-guard")) + .await + .unwrap(); let mut pending = Box::pin(handle_update_config(&state, authed_request(update))); // Poll the actual writer while an endpoint report owns the mutation @@ -442,6 +450,39 @@ async fn global_policy_update_waits_for_endpoint_report_guard() { } } +#[tokio::test] +async fn report_endpoint_status_does_not_wait_for_unrelated_sandbox_guard() { + let sandbox_id = "endpoint-unrelated-guard"; + let (state, mut report) = sandbox_with_accepted_endpoint_result(sandbox_id, true).await; + let unrelated_guard = state + .compute + .mutation_guard(MutationScope::sandbox( + "default", + "endpoint-unrelated-other", + )) + .await + .expect("hold an unrelated sandbox guard"); + + report.report_sequence = 2; + report.observations[0].result = EndpointResult::TransportFailed as i32; + tokio::time::timeout( + std::time::Duration::from_secs(5), + handle_report_endpoint_status(&state, with_sandbox(Request::new(report), sandbox_id)), + ) + .await + .expect("an endpoint report must not wait for an unrelated sandbox mutation") + .expect("accept endpoint result"); + let status = stored_sandbox(&state, sandbox_id) + .await + .status + .expect("sandbox status"); + assert_eq!( + status.endpoint_statuses[0].last_result, + EndpointResult::TransportFailed as i32 + ); + drop(unrelated_guard); +} + #[test] fn expected_endpoint_statuses_canonicalize_identity_and_distinguish_paths() { let policy = ProtoSandboxPolicy { @@ -836,6 +877,583 @@ async fn startup_reconciliation_invalidates_status_from_previous_sessions() { assert_eq!(status.conditions, vec![ready_condition()]); } +/// Store a sandbox carrying endpoint evidence from an earlier gateway process +/// and return the status startup reconciliation must leave behind. +async fn seed_stale_endpoint_status(state: &ServerState, sandbox_id: &str) -> EndpointStatus { + let mut sandbox = test_sandbox( + sandbox_id, + sandbox_id, + mcp_policy_with_versions(&["2025-11-25"]), + Vec::new(), + ); + let initial = test_initial_endpoint_status(sandbox_id, "api.example.com", "/mcp"); + sandbox.status = Some(SandboxStatus { + endpoint_statuses: vec![EndpointStatus { + last_result: EndpointResult::HttpResponseReceived as i32, + last_reported_time: Some(timestamp("2026-09-05T01:01:00.000Z")), + ..initial.clone() + }], + conditions: vec![ready_condition()], + ..Default::default() + }); + state + .store + .put_message(&sandbox) + .await + .expect("store prior session status"); + initial +} + +fn spawn_startup_reconciliation( + state: &Arc, +) -> tokio::task::JoinHandle> { + let state = state.clone(); + tokio::spawn(async move { invalidate_endpoint_status_on_startup(&state).await }) +} + +#[tokio::test] +async fn startup_reconciliation_does_not_hold_a_fleet_guard() { + let state = test_server_state().await; + let sandbox_id = "endpoint-startup-fleet-target"; + let initial = seed_stale_endpoint_status(&state, sandbox_id).await; + let unrelated_guard = state + .compute + .mutation_guard(MutationScope::sandbox( + "default", + "endpoint-startup-fleet-unrelated", + )) + .await + .expect("hold an unrelated sandbox guard"); + + tokio::time::timeout( + std::time::Duration::from_secs(5), + invalidate_endpoint_status_on_startup(&state), + ) + .await + .expect("startup reconciliation must not wait for an unrelated sandbox mutation") + .expect("startup reconciliation"); + let status = stored_sandbox(&state, sandbox_id) + .await + .status + .expect("status remains present"); + assert_eq!(status.endpoint_statuses, vec![initial]); + drop(unrelated_guard); +} + +#[tokio::test] +async fn startup_reconciliation_skips_sandbox_deleted_while_waiting() { + let state = test_server_state().await; + let sandbox_id = "endpoint-startup-deleted"; + seed_stale_endpoint_status(&state, sandbox_id).await; + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", sandbox_id)) + .await + .expect("hold the target sandbox guard"); + let mut reconciliation = spawn_startup_reconciliation(&state); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), &mut reconciliation) + .await + .is_err(), + "startup reconciliation must wait for the target sandbox guard" + ); + + // The sandbox was listed before this delete, so the reconciliation that + // wakes up must treat the missing row as done rather than fail startup. + assert!( + state + .store + .delete("sandbox", sandbox_id) + .await + .expect("delete sandbox") + ); + drop(guard); + tokio::time::timeout(std::time::Duration::from_secs(5), reconciliation) + .await + .expect("startup reconciliation finishes after the guard is released") + .expect("startup reconciliation task") + .expect("a sandbox deleted while reconciliation waited is skipped"); +} + +#[tokio::test] +async fn startup_reconciliation_pages_past_a_sandbox_deleted_mid_scan() { + let state = test_server_state().await; + let page_size = usize::try_from(ENDPOINT_STARTUP_RECONCILIATION_PAGE_SIZE).unwrap(); + // Zero-padded ids keep the listing order equal to the seeding order. + let ids: Vec = (0..page_size + 2) + .map(|index| format!("endpoint-startup-page-{index:03}")) + .collect(); + let mut initial = Vec::with_capacity(ids.len()); + for sandbox_id in &ids { + initial.push(seed_stale_endpoint_status(&state, sandbox_id).await); + } + let is_reset = |sandbox_id: &str, expected: &EndpointStatus| { + let state = state.clone(); + let sandbox_id = sandbox_id.to_string(); + let expected = expected.clone(); + async move { + stored_sandbox(&state, &sandbox_id) + .await + .status + .is_some_and(|status| status.endpoint_statuses == vec![expected]) + } + }; + + // Holding the first sandbox lets the scan reset the rest of page 1, then + // keeps it from reading page 2. + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", &ids[0])) + .await + .expect("hold the first sandbox guard"); + let reconciliation = spawn_startup_reconciliation(&state); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + for (sandbox_id, expected) in ids[1..page_size].iter().zip(&initial[1..page_size]) { + while !is_reset(sandbox_id, expected).await { + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + } + } + }) + .await + .expect("the scan resets the rest of page 1 while it waits"); + assert!(!reconciliation.is_finished()); + + // With offset paging, page 2 would now skip its first sandbox. + let deleted = page_size / 2; + assert!( + state + .store + .delete("sandbox", &ids[deleted]) + .await + .expect("delete sandbox") + ); + drop(guard); + tokio::time::timeout(std::time::Duration::from_secs(5), reconciliation) + .await + .expect("startup reconciliation finishes after the guard is released") + .expect("startup reconciliation task") + .expect("startup reconciliation"); + for (index, (sandbox_id, expected)) in ids.iter().zip(&initial).enumerate() { + if index != deleted { + assert!( + is_reset(sandbox_id, expected).await, + "{sandbox_id} kept stale endpoint status" + ); + } + } +} + +#[tokio::test] +async fn startup_reconciliation_rereads_under_guard() { + let state = test_server_state().await; + let sandbox_id = "endpoint-startup-reread"; + let initial = seed_stale_endpoint_status(&state, sandbox_id).await; + let listed_version = stored_sandbox(&state, sandbox_id) + .await + .get_resource_version(); + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", sandbox_id)) + .await + .expect("hold the target sandbox guard"); + let mut reconciliation = spawn_startup_reconciliation(&state); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), &mut reconciliation) + .await + .is_err(), + "startup reconciliation must wait for the target sandbox guard" + ); + + // Lifecycle writers on other replicas take no distributed guard, so the + // listed version can go stale while reconciliation waits. + let bumped = state + .store + .update_message_cas::(sandbox_id, listed_version, |sandbox| { + sandbox + .metadata + .as_mut() + .expect("sandbox metadata") + .labels + .insert("concurrent-write".to_string(), "true".to_string()); + }) + .await + .expect("concurrent unrelated write"); + assert!(bumped.get_resource_version() > listed_version); + drop(guard); + tokio::time::timeout(std::time::Duration::from_secs(5), reconciliation) + .await + .expect("startup reconciliation finishes after the guard is released") + .expect("startup reconciliation task") + .expect("startup reconciliation"); + + let sandbox = stored_sandbox(&state, sandbox_id).await; + assert!(sandbox.get_resource_version() > bumped.get_resource_version()); + assert_eq!( + sandbox + .metadata + .as_ref() + .expect("sandbox metadata") + .labels + .get("concurrent-write") + .map(String::as_str), + Some("true") + ); + assert_eq!( + sandbox + .status + .expect("status remains present") + .endpoint_statuses, + vec![initial] + ); +} + +/// Assert that the evidence `seed_stale_endpoint_status` stored survived. +async fn assert_endpoint_evidence_kept(state: &ServerState, sandbox_id: &str) { + let status = stored_sandbox(state, sandbox_id) + .await + .status + .expect("status remains present"); + assert_eq!( + status.endpoint_statuses[0].last_result, + EndpointResult::HttpResponseReceived as i32 + ); + assert!(status.endpoint_statuses[0].last_reported_time.is_some()); +} + +#[tokio::test] +async fn startup_reconciliation_keeps_evidence_of_a_live_owner() { + let state = test_server_state().await; + let sandbox_id = "endpoint-startup-live-owner"; + seed_stale_endpoint_status(&state, sandbox_id).await; + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", sandbox_id)) + .await + .expect("hold the target sandbox guard"); + let mut reconciliation = spawn_startup_reconciliation(&state); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), &mut reconciliation) + .await + .is_err(), + "startup reconciliation must wait for the target sandbox guard" + ); + + // The supervisor connects to a peer after the unguarded owner check, so + // only the re-check under the guard can see it. + SupervisorOwnerIndex::new(state.store.clone(), OWNER_TTL) + .publish( + sandbox_id, + "session", + "supervisor", + 1, + "peer-replica", + "https://peer", + ) + .await + .expect("publish a live owner on a peer"); + drop(guard); + tokio::time::timeout(std::time::Duration::from_secs(5), reconciliation) + .await + .expect("startup reconciliation finishes after the guard is released") + .expect("startup reconciliation task") + .expect("startup reconciliation"); + assert_endpoint_evidence_kept(&state, sandbox_id).await; + + // With the owner already live, the check before the guard skips the + // sandbox without waiting for its mutation. + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", sandbox_id)) + .await + .expect("hold the target sandbox guard"); + tokio::time::timeout( + std::time::Duration::from_secs(5), + invalidate_endpoint_status_on_startup(&state), + ) + .await + .expect("a live owner is skipped without waiting for the sandbox guard") + .expect("startup reconciliation"); + drop(guard); + assert_endpoint_evidence_kept(&state, sandbox_id).await; +} + +#[tokio::test] +async fn startup_reconciliation_retries_a_conflict_then_succeeds() { + let mut calls = 0; + retry_startup_reconciliation(|| { + calls += 1; + let call = calls; + async move { + Ok(if call == 1 { + StartupAttempt::Retry(PersistenceError::Conflict { + current_resource_version: Some(2), + }) + } else { + StartupAttempt::Done + }) + } + }) + .await + .expect("a conflict is retried"); + assert_eq!(calls, 2); +} + +#[tokio::test] +async fn startup_reconciliation_stops_after_the_attempt_limit() { + let mut calls = 0; + let error = retry_startup_reconciliation(|| { + calls += 1; + async { + Ok(StartupAttempt::Retry(PersistenceError::Database( + "object sb not found".to_string(), + ))) + } + }) + .await + .expect_err("retries are bounded"); + assert_eq!(error.code(), Code::Internal); + assert_eq!(calls, ENDPOINT_STARTUP_RECONCILIATION_ATTEMPTS); + + let mut calls = 0; + let error = retry_startup_reconciliation(|| { + calls += 1; + async { Err(Status::unavailable("resolve supervisor owner failed")) } + }) + .await + .expect_err("an attempt error is returned"); + assert_eq!(error.code(), Code::Unavailable); + assert_eq!(calls, 1, "an attempt error is not retried"); +} + +/// Store accepted endpoint evidence, then end the supervisor session that +/// reported it, as the session task does before its disconnect reset. +async fn disconnected_sandbox_with_endpoint_result(sandbox_id: &str) -> Arc { + let (state, _report) = sandbox_with_accepted_endpoint_result(sandbox_id, true).await; + assert!( + state + .supervisor_sessions + .remove_if_current(sandbox_id, "session-a") + .is_some() + ); + assert_eq!( + stored_sandbox(&state, sandbox_id) + .await + .status + .expect("sandbox status") + .endpoint_statuses[0] + .last_result, + EndpointResult::HttpResponseReceived as i32 + ); + state +} + +#[tokio::test] +async fn disconnect_reset_skips_guard_when_local_session_replaced() { + let sandbox_id = "endpoint-disconnect-local-replacement"; + let state = disconnected_sandbox_with_endpoint_result(sandbox_id).await; + register_session(&state, sandbox_id, "session-b"); + let before = stored_sandbox(&state, sandbox_id).await; + let guard = state + .compute + .mutation_guard(MutationScope::Global) + .await + .expect("hold a guard that excludes every sandbox guard"); + + tokio::time::timeout( + std::time::Duration::from_secs(5), + reset_endpoint_status_after_supervisor_disconnect(&state, sandbox_id), + ) + .await + .expect("a disconnect with a local replacement must not wait for the guard") + .expect("disconnect reset"); + assert_eq!(stored_sandbox(&state, sandbox_id).await, before); + drop(guard); +} + +#[tokio::test] +async fn disconnect_reset_skips_guard_when_peer_owns_session() { + let sandbox_id = "endpoint-disconnect-peer-owner"; + let state = disconnected_sandbox_with_endpoint_result(sandbox_id).await; + SupervisorOwnerIndex::new(state.store.clone(), OWNER_TTL) + .publish( + sandbox_id, + "peer-s", + "inst", + 1, + "peer-replica", + "https://peer:8080", + ) + .await + .expect("publish a live owner on a peer"); + let before = stored_sandbox(&state, sandbox_id).await; + let guard = state + .compute + .mutation_guard(MutationScope::Global) + .await + .expect("hold a guard that excludes every sandbox guard"); + + tokio::time::timeout( + std::time::Duration::from_secs(5), + reset_endpoint_status_after_supervisor_disconnect(&state, sandbox_id), + ) + .await + .expect("a disconnect with a live peer owner must not wait for the guard") + .expect("disconnect reset"); + assert_eq!(stored_sandbox(&state, sandbox_id).await, before); + drop(guard); +} + +#[tokio::test] +async fn disconnect_reset_waits_for_guard_without_replacement() { + let sandbox_id = "endpoint-disconnect-no-replacement"; + let state = disconnected_sandbox_with_endpoint_result(sandbox_id).await; + let before = stored_sandbox(&state, sandbox_id).await; + let guard = state + .compute + .mutation_guard(MutationScope::Global) + .await + .expect("hold a guard that excludes every sandbox guard"); + let mut pending = Box::pin(reset_endpoint_status_after_supervisor_disconnect( + &state, sandbox_id, + )); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), pending.as_mut()) + .await + .is_err(), + "a disconnect without a replacement must wait for the guard" + ); + assert_eq!(stored_sandbox(&state, sandbox_id).await, before); + + drop(guard); + tokio::time::timeout(std::time::Duration::from_secs(5), pending) + .await + .expect("disconnect reset finishes after the guard is released") + .expect("disconnect reset"); + let status = stored_sandbox(&state, sandbox_id) + .await + .status + .expect("status remains present"); + assert_eq!( + status.endpoint_statuses[0].last_result, + EndpointResult::NoObservedExchange as i32 + ); + assert!(status.endpoint_statuses[0].last_reported_time.is_none()); +} + +#[tokio::test] +async fn disconnect_reset_rechecks_for_replacement_under_guard() { + let sandbox_id = "endpoint-disconnect-late-replacement"; + let state = disconnected_sandbox_with_endpoint_result(sandbox_id).await; + let guard = state + .compute + .mutation_guard(MutationScope::Global) + .await + .expect("hold a guard that excludes every sandbox guard"); + let mut pending = Box::pin(reset_endpoint_status_after_supervisor_disconnect( + &state, sandbox_id, + )); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), pending.as_mut()) + .await + .is_err(), + "a disconnect without a replacement must wait for the guard" + ); + + // The replacement registers after the unguarded check, so only the + // re-check under the guard can keep its evidence. + register_session(&state, sandbox_id, "session-b"); + let before = stored_sandbox(&state, sandbox_id).await; + drop(guard); + tokio::time::timeout(std::time::Duration::from_secs(5), pending) + .await + .expect("disconnect reset finishes after the guard is released") + .expect("disconnect reset"); + assert_eq!(stored_sandbox(&state, sandbox_id).await, before); +} + +/// Serve the gateway API on loopback so a test can open a real supervisor +/// stream. +async fn gateway_client(state: Arc) -> OpenShellClient { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind loopback listener"); + let address = listener.local_addr().expect("listener address"); + tokio::spawn( + tonic::transport::Server::builder() + .add_service(OpenShellServer::new(OpenShellService::new(state))) + .serve_with_incoming(TcpListenerStream::new(listener)), + ); + OpenShellClient::connect(format!("http://{address}")) + .await + .expect("connect to the test gateway") +} + +#[tokio::test] +async fn replacement_reset_timeout_invalidates_superseded_evidence() { + let sandbox_id = "endpoint-replacement-reset-timeout"; + // Session A stays registered, so the replacement supersedes it before A's + // session loop could schedule a disconnect reset. + let (state, _report) = sandbox_with_accepted_endpoint_result(sandbox_id, true).await; + let hold = state + .compute + .sandbox_mutation_guard_by_id(sandbox_id) + .await + .expect("hold the sandbox mutation guard") + .expect("sandbox exists"); + state + .compute + .set_mutation_lock_timeout_for_tests(std::time::Duration::from_millis(50)); + + let hello = SupervisorMessage { + payload: Some(supervisor_message::Payload::Hello(SupervisorHello { + sandbox_id: sandbox_id.to_string(), + instance_id: "instance-b".to_string(), + connection_epoch: 1, + supports_provider_readiness: false, + })), + }; + let error = gateway_client(Arc::clone(&state)) + .await + .connect_supervisor(tokio_stream::iter([hello])) + .await + .expect_err("the replacement's endpoint reset times out on the held guard"); + assert_eq!(error.code(), Code::Unavailable); + assert!( + state + .supervisor_sessions + .current_session_id(sandbox_id) + .is_none() + ); + assert_eq!( + stored_sandbox(&state, sandbox_id) + .await + .status + .expect("sandbox status") + .endpoint_statuses[0] + .last_result, + EndpointResult::HttpResponseReceived as i32, + "nothing resets while the guard is held" + ); + + drop(hold); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let status = stored_sandbox(&state, sandbox_id) + .await + .status + .expect("sandbox status"); + if status.endpoint_statuses[0].last_result == EndpointResult::NoObservedExchange as i32 + { + assert!(status.endpoint_statuses[0].last_reported_time.is_none()); + return; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + }) + .await + .expect("the superseded session's evidence is reset once the guard is free"); +} + #[tokio::test] async fn report_endpoint_status_is_session_bound_and_retry_idempotent() { let state = test_server_state().await; diff --git a/crates/openshell-server/src/grpc/provider.rs b/crates/openshell-server/src/grpc/provider.rs index fb22dc698d..9860c9ffe2 100644 --- a/crates/openshell-server/src/grpc/provider.rs +++ b/crates/openshell-server/src/grpc/provider.rs @@ -5,6 +5,7 @@ #![allow(clippy::result_large_err)] // gRPC handlers return Result, Status> +use crate::compute::MutationScope; #[cfg(test)] use crate::credentials::RefreshMaterialScope; use crate::pagination::Pagination; @@ -2609,9 +2610,13 @@ pub(super) async fn handle_create_provider( )); } let provider_type = provider.r#type.clone(); - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::persistence_error_to_status(error, "acquire provider mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .map_err(|error| { + super::persistence_error_to_status(error, "acquire provider mutation lock") + })?; let catalog = state .provider_profile_sources .snapshot_catalog(state.store.as_ref(), &workspace) @@ -2838,10 +2843,11 @@ pub(super) async fn handle_import_provider_profiles( .ensure_active()?; let (profiles, mut diagnostics) = profiles_from_import_items(&request.profiles); add_empty_profile_set_diagnostic(&profiles, &mut diagnostics); - let _sandbox_sync_guard = - state.compute.sandbox_sync_guard().await.map_err(|err| { - super::persistence_error_to_status(err, "acquire provider mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .map_err(|err| super::persistence_error_to_status(err, "acquire provider mutation lock"))?; let catalog = state .provider_profile_sources .snapshot_catalog(state.store.as_ref(), &workspace) @@ -2933,10 +2939,11 @@ pub(super) async fn handle_update_provider_profiles( let (profiles, mut diagnostics) = profiles_from_import_items(&items); add_empty_profile_set_diagnostic(&profiles, &mut diagnostics); let target_id = normalize_profile_id_request(&request.id)?; - let _sandbox_sync_guard = - state.compute.sandbox_sync_guard().await.map_err(|err| { - super::persistence_error_to_status(err, "acquire provider mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .map_err(|err| super::persistence_error_to_status(err, "acquire provider mutation lock"))?; let catalog = state .provider_profile_sources .snapshot_catalog(state.store.as_ref(), &workspace) @@ -3096,10 +3103,11 @@ pub(super) async fn handle_delete_provider_profile( .name; let id = req.id; let id = normalize_profile_id_request(&id)?; - let _sandbox_sync_guard = - state.compute.sandbox_sync_guard().await.map_err(|err| { - super::persistence_error_to_status(err, "acquire provider mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .map_err(|err| super::persistence_error_to_status(err, "acquire provider mutation lock"))?; let catalog = state .provider_profile_sources .snapshot_catalog(state.store.as_ref(), &workspace) @@ -3922,10 +3930,15 @@ pub(super) async fn handle_update_provider( .name; // Provider material contributes to the route-report configuration epoch. // Serialize its mutation with route-status validation so a report derived - // from the prior revision cannot commit after this update. - let _sandbox_sync_guard = state.compute.sandbox_sync_guard().await.map_err(|error| { - super::persistence_error_to_status(error, "acquire provider mutation lock") - })?; + // from the prior revision cannot commit after this update. The workspace + // key excludes every sandbox mutation in this workspace. + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .map_err(|error| { + super::persistence_error_to_status(error, "acquire provider mutation lock") + })?; let Some(mut provider) = req.provider else { emit_provider_lifecycle( "custom", @@ -4680,11 +4693,14 @@ pub(super) async fn handle_configure_provider_refresh( // persist further down are otherwise separate steps: two concurrent // configures of providers attached to the same sandbox could each pass // validation before either persisted and both reserve the same key (CWE-362). - // This is the same guard sandbox create/attach and profile changes take. - let _sandbox_sync_guard = - state.compute.sandbox_sync_guard().await.map_err(|err| { - super::persistence_error_to_status(err, "acquire provider mutation lock") - })?; + // Every provider attached to one sandbox lives in this workspace, so holding + // the workspace key exclusively also excludes sandbox create and attach, + // which hold it shared. + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .map_err(|err| super::persistence_error_to_status(err, "acquire provider mutation lock"))?; let provider = state .store @@ -5125,9 +5141,22 @@ pub(super) async fn handle_delete_provider( MinWorkspaceRole::Admin, ) .await?; + // Reject after authorization but before taking the workspace lock, which + // a request that can never succeed should not wait for. + if req.name.is_empty() { + return Err(Status::invalid_argument("name is required")); + } let workspace = super::workspace::resolve_workspace(state.store.as_ref(), &authz.workspace) .await? .name; + // A sandbox create or attach in this workspace holds the workspace key + // shared, so no sandbox can start referencing the provider between the + // attached-sandbox check and the delete. + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::Workspace(&workspace)) + .await + .map_err(|e| super::persistence_error_to_status(e, "acquire provider mutation lock"))?; let name = req.name; let provider_profile = provider_profile_for_name(state.store.as_ref(), &workspace, &name).await; let result = delete_provider_record_with_credentials( @@ -5760,9 +5789,13 @@ mod tests { } #[tokio::test] - async fn import_provider_profile_waits_for_sandbox_sync_guard() { + async fn import_provider_profile_waits_for_sandbox_mutation_in_workspace() { let state = test_server_state().await; - let guard = state.compute.sandbox_sync_guard().await.unwrap(); + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", "profile-import-guard")) + .await + .unwrap(); let task_state = state.clone(); let task = tokio::spawn(async move { handle_import_provider_profiles( @@ -5784,7 +5817,7 @@ mod tests { tokio::time::sleep(std::time::Duration::from_millis(20)).await; assert!( !task.is_finished(), - "profile import should wait for sandbox sync guard" + "profile import should wait for a sandbox mutation in its workspace" ); drop(guard); @@ -8939,7 +8972,7 @@ mod tests { } #[tokio::test] - async fn delete_provider_profile_waits_for_sandbox_sync_guard() { + async fn delete_provider_profile_waits_for_sandbox_mutation_in_workspace() { let state = test_server_state().await; state .store @@ -8947,7 +8980,11 @@ mod tests { .await .unwrap(); - let guard = state.compute.sandbox_sync_guard().await.unwrap(); + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", "profile-delete-guard")) + .await + .unwrap(); let task_state = state.clone(); let task = tokio::spawn(async move { handle_delete_provider_profile( @@ -8967,7 +9004,7 @@ mod tests { tokio::time::sleep(std::time::Duration::from_millis(20)).await; assert!( !task.is_finished(), - "profile delete should wait for sandbox sync guard" + "profile delete should wait for a sandbox mutation in its workspace" ); drop(guard); @@ -9002,7 +9039,11 @@ mod tests { .await .unwrap(); - let guard = state.compute.sandbox_sync_guard().await.unwrap(); + let guard = state + .compute + .mutation_guard(MutationScope::Workspace("default")) + .await + .unwrap(); let task_state = state.clone(); let task = tokio::spawn(async move { let mut provider = provider_with_values("guarded-provider", "guarded-create"); @@ -9040,6 +9081,232 @@ mod tests { ); } + fn default_workspace_selector() -> openshell_core::proto::WorkspaceSelector { + openshell_core::proto::workspace_selector("default".to_string()) + } + + async fn create_openai_provider(state: &Arc, name: &str) -> Provider { + let provider = provider_with_credential_value(name, "openai", "OPENAI_API_KEY", "sk-test"); + handle_create_provider( + state, + authed_request(CreateProviderRequest { + request_id: String::new(), + provider: Some(provider), + workspace_scope: Some(default_workspace_selector()), + }), + ) + .await + .expect("create provider") + .into_inner() + .provider + .expect("created provider") + } + + fn provider_config_update(current: &Provider) -> Request { + let mut provider = current.clone(); + provider.credential_handles.clear(); + provider + .config + .insert("NEW_CONFIG".to_string(), "new-value".to_string()); + authed_request(UpdateProviderRequest { + request_id: String::new(), + provider: Some(provider), + credential_expiration_times: HashMap::new(), + clear_credential_expiration_keys: Vec::new(), + workspace_scope: Some(default_workspace_selector()), + }) + } + + fn delete_provider_request(name: &str) -> Request { + authed_request(DeleteProviderRequest { + request_id: String::new(), + allow_missing: false, + name: name.to_string(), + workspace_scope: Some(default_workspace_selector()), + }) + } + + fn sandbox_in_default_workspace(id: &str, providers: Vec) -> Sandbox { + Sandbox { + metadata: Some(openshell_core::proto::datamodel::v1::ObjectMeta { + id: id.to_string(), + name: id.to_string(), + workspace: "default".to_string(), + ..Default::default() + }), + spec: Some(SandboxSpec { + providers, + ..Default::default() + }), + ..Default::default() + } + } + + #[tokio::test] + async fn delete_provider_waits_for_sandbox_mutation_in_workspace() { + let state = test_server_state().await; + create_openai_provider(&state, "guarded-delete-provider").await; + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", "x")) + .await + .unwrap(); + + let task_state = state.clone(); + let mut delete = tokio::spawn(async move { + handle_delete_provider( + &task_state, + delete_provider_request("guarded-delete-provider"), + ) + .await + }); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), &mut delete) + .await + .is_err(), + "provider delete should wait for a sandbox mutation in its workspace" + ); + drop(guard); + + let response = tokio::time::timeout(std::time::Duration::from_secs(5), delete) + .await + .expect("delete should finish after guard release") + .expect("join delete task") + .expect("delete should succeed") + .into_inner(); + assert_eq!( + response.outcome(), + openshell_core::proto::DeletionOutcome::Completed + ); + } + + #[tokio::test] + async fn delete_provider_rejects_empty_name_without_waiting_for_workspace() { + let state = test_server_state().await; + let _guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", "x")) + .await + .unwrap(); + + let error = tokio::time::timeout( + std::time::Duration::from_secs(1), + handle_delete_provider(&state, delete_provider_request("")), + ) + .await + .expect("an empty name should not wait for the workspace lock") + .unwrap_err(); + assert_eq!(error.code(), Code::InvalidArgument); + } + + #[tokio::test] + async fn delete_provider_rejects_provider_attached_while_waiting() { + let state = test_server_state().await; + create_openai_provider(&state, "raced-provider").await; + let sandbox = sandbox_in_default_workspace("raced-sandbox", Vec::new()); + state.store.put_message(&sandbox).await.unwrap(); + // An attach to this sandbox holds its sandbox scope while it writes + // the provider into the spec. + let attach_guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", sandbox.object_id())) + .await + .unwrap(); + + let task_state = state.clone(); + let mut delete = tokio::spawn(async move { + handle_delete_provider(&task_state, delete_provider_request("raced-provider")).await + }); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), &mut delete) + .await + .is_err(), + "provider delete should wait for the in-flight attach" + ); + state + .store + .update_message_cas::(sandbox.object_id(), 0, |sandbox| { + sandbox + .spec + .get_or_insert_with(Default::default) + .providers + .push("raced-provider".to_string()); + }) + .await + .unwrap(); + drop(attach_guard); + + let error = tokio::time::timeout(std::time::Duration::from_secs(5), delete) + .await + .expect("delete should finish after the attach") + .expect("join delete task") + .expect_err("a provider attached while the delete waited must not be deleted"); + assert_eq!(error.code(), Code::FailedPrecondition); + assert!(error.message().contains("attached to sandbox"), "{error}"); + assert!( + state + .store + .get_message_by_name::("default", "raced-provider") + .await + .unwrap() + .is_some() + ); + } + + #[tokio::test] + async fn update_provider_waits_for_sandbox_mutation_in_same_workspace() { + let state = test_server_state().await; + let current = create_openai_provider(&state, "guarded-update-provider").await; + let guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", "x")) + .await + .unwrap(); + + let task_state = state.clone(); + let request = provider_config_update(¤t); + let mut update = + tokio::spawn(async move { handle_update_provider(&task_state, request).await }); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), &mut update) + .await + .is_err(), + "provider update should wait for a sandbox mutation in its workspace" + ); + drop(guard); + + let response = tokio::time::timeout(std::time::Duration::from_secs(5), update) + .await + .expect("update should finish after guard release") + .expect("join update task") + .expect("update should succeed") + .into_inner(); + assert!(response.provider.unwrap().config.contains_key("NEW_CONFIG")); + } + + #[tokio::test] + async fn update_provider_does_not_wait_for_sandbox_mutation_in_other_workspace() { + let state = test_server_state().await; + let current = create_openai_provider(&state, "unguarded-update-provider").await; + // The "team-a" workspace row is not needed to hold its keys. + let other_workspace_guard = state + .compute + .mutation_guard(MutationScope::sandbox("team-a", "x")) + .await + .unwrap(); + + let response = tokio::time::timeout( + std::time::Duration::from_secs(5), + handle_update_provider(&state, provider_config_update(¤t)), + ) + .await + .expect("provider update should not wait for another workspace") + .expect("update should succeed") + .into_inner(); + assert!(response.provider.unwrap().config.contains_key("NEW_CONFIG")); + drop(other_workspace_guard); + } + #[tokio::test] async fn provider_crud_round_trip_and_semantics() { let store = test_store().await; diff --git a/crates/openshell-server/src/grpc/provider_readiness_tests.rs b/crates/openshell-server/src/grpc/provider_readiness_tests.rs index d83bb80ef4..4e82b3bba0 100644 --- a/crates/openshell-server/src/grpc/provider_readiness_tests.rs +++ b/crates/openshell-server/src/grpc/provider_readiness_tests.rs @@ -824,7 +824,7 @@ async fn attach_waiting_for_update_captures_published_revision_and_becomes_ready .provider .unwrap(); - // The credential driver's gate holds UpdateProvider inside the shared + // The credential driver's gate holds UpdateProvider inside the workspace // mutation guard while the attach request reaches that same guard. let (store_hit, release_store) = state.credentials.gate_next_store(); let update_state = Arc::clone(&state); diff --git a/crates/openshell-server/src/grpc/sandbox.rs b/crates/openshell-server/src/grpc/sandbox.rs index f1c67be49d..db0a096a6f 100644 --- a/crates/openshell-server/src/grpc/sandbox.rs +++ b/crates/openshell-server/src/grpc/sandbox.rs @@ -14,6 +14,7 @@ use crate::auth::workspace_authz::{ AuthorizedWorkspaceScope, MinWorkspaceRole, authorize_list_workspace_selector, authorize_sandbox_workspace, authorize_workspace, }; +use crate::compute::MutationScope; use crate::pagination::Pagination; use crate::persistence::{ ObjectLabels, ObjectListQuery, ObjectType, WriteCondition, generate_name, @@ -507,9 +508,9 @@ async fn handle_create_sandbox_inner( } else { request.name.clone() }; - let (sandbox_lifecycle_guard, sandbox_sync_guard) = state + let (sandbox_lifecycle_guard, mutation_guard) = state .compute - .sandbox_create_guards(&id) + .sandbox_create_guards(&workspace, &id) .await .map_err(|err| super::persistence_error_to_status(err, "acquire sandbox mutation lock"))?; @@ -677,7 +678,7 @@ async fn handle_create_sandbox_inner( launch_authentication, await_main_process_attachment, sandbox_lifecycle_guard, - sandbox_sync_guard, + mutation_guard, )) .await?; @@ -1334,10 +1335,14 @@ pub(super) async fn handle_attach_sandbox_provider( if let Some(probe) = attach_wait_probe { probe.notify_one(); } - let _sandbox_sync_guard = - state.compute.sandbox_sync_guard().await.map_err(|err| { - super::persistence_error_to_status(err, "acquire sandbox mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::sandbox( + sandbox.object_workspace(), + sandbox.object_id(), + )) + .await + .map_err(|err| super::persistence_error_to_status(err, "acquire sandbox mutation lock"))?; let provider_record = get_provider_record(state.store.as_ref(), &workspace, &request.provider) .await .map_err(|err| { @@ -1506,10 +1511,14 @@ pub(super) async fn handle_detach_sandbox_provider( ))); } - let _sandbox_sync_guard = - state.compute.sandbox_sync_guard().await.map_err(|err| { - super::persistence_error_to_status(err, "acquire sandbox mutation lock") - })?; + let _mutation_guard = state + .compute + .mutation_guard(MutationScope::sandbox( + sandbox.object_workspace(), + sandbox.object_id(), + )) + .await + .map_err(|err| super::persistence_error_to_status(err, "acquire sandbox mutation lock"))?; let sandbox_name = sandbox.object_name().to_string(); let sandbox_id = sandbox .metadata @@ -5575,8 +5584,13 @@ mod tests { state.store.put_message(&original).await.unwrap(); // Hold the global guard so the handler can resolve the original ID and - // acquire its delete gate, but cannot yet revalidate or mutate it. - let global_guard = state.compute.sandbox_sync_guard().await.unwrap(); + // acquire its delete gate, but cannot yet take the sandbox's local + // lifecycle lock (shared global key) to revalidate or mutate it. + let global_guard = state + .compute + .mutation_guard(MutationScope::Global) + .await + .unwrap(); let delete_state = state.clone(); let delete = tokio::spawn(async move { handle_delete_sandbox_inner( @@ -5757,6 +5771,91 @@ mod tests { assert_eq!(providers, vec!["work-github"]); } + fn attach_request(sandbox: &str, provider: &str) -> Request { + authed_request(AttachSandboxProviderRequest { + request_id: String::new(), + sandbox: sandbox.to_string(), + workspace_scope: Some(openshell_core::proto::workspace_selector( + "default".to_string(), + )), + provider: provider.to_string(), + expected_resource_version: 0, + }) + } + + #[tokio::test] + async fn attach_provider_does_not_wait_for_unrelated_sandbox_guard() { + let state = test_server_state().await; + state + .store + .put_message(&test_provider("work-github", "github")) + .await + .unwrap(); + state + .store + .put_message(&test_sandbox("work", Vec::new())) + .await + .unwrap(); + let unrelated = test_sandbox("unrelated", Vec::new()); + state.store.put_message(&unrelated).await.unwrap(); + let unrelated_guard = state + .compute + .mutation_guard(MutationScope::sandbox("default", unrelated.object_id())) + .await + .unwrap(); + + let response = tokio::time::timeout( + std::time::Duration::from_secs(5), + handle_attach_sandbox_provider(&state, attach_request("work", "work-github")), + ) + .await + .expect("attach should not wait for an unrelated sandbox mutation") + .expect("attach should succeed") + .into_inner(); + assert!(response.attached); + drop(unrelated_guard); + } + + #[tokio::test] + async fn attach_provider_waits_for_workspace_guard() { + let state = test_server_state().await; + state + .store + .put_message(&test_provider("work-github", "github")) + .await + .unwrap(); + state + .store + .put_message(&test_sandbox("work", Vec::new())) + .await + .unwrap(); + let workspace_guard = state + .compute + .mutation_guard(MutationScope::Workspace("default")) + .await + .unwrap(); + + let task_state = state.clone(); + let mut attach = tokio::spawn(async move { + handle_attach_sandbox_provider(&task_state, attach_request("work", "work-github")).await + }); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), &mut attach) + .await + .is_err(), + "attach should wait for a provider writer in its workspace" + ); + drop(workspace_guard); + + let response = tokio::time::timeout(std::time::Duration::from_secs(5), attach) + .await + .expect("attach should finish after the workspace guard is released") + .expect("join attach task") + .expect("attach should succeed") + .into_inner(); + assert!(response.attached); + } + #[tokio::test] async fn detach_sandbox_provider_is_idempotent_and_removes_all_matches() { let state = test_server_state().await; @@ -7222,7 +7321,7 @@ mod tests { } #[tokio::test] - async fn create_sandbox_with_providers_waits_for_sandbox_sync_guard() { + async fn create_sandbox_with_providers_waits_for_workspace_mutation_guard() { let state = test_server_state().await; state .store @@ -7230,7 +7329,13 @@ mod tests { .await .unwrap(); - let guard = state.compute.sandbox_sync_guard().await.unwrap(); + // A provider writer in the workspace excludes creates that validate + // against its providers. + let guard = state + .compute + .mutation_guard(MutationScope::Workspace("default")) + .await + .unwrap(); let task_state = state.clone(); let task = tokio::spawn(async move { handle_create_sandbox( @@ -7258,7 +7363,7 @@ mod tests { tokio::time::sleep(std::time::Duration::from_millis(20)).await; assert!( !task.is_finished(), - "sandbox create with initial providers should wait for sandbox sync guard" + "sandbox create with initial providers should wait for the workspace mutation guard" ); drop(guard); diff --git a/crates/openshell-server/src/lib.rs b/crates/openshell-server/src/lib.rs index 7ec44d084a..fb111079d2 100644 --- a/crates/openshell-server/src/lib.rs +++ b/crates/openshell-server/src/lib.rs @@ -288,12 +288,6 @@ pub struct ServerState { /// Active SSH tunnel connection counts per sandbox id. pub ssh_connections_by_sandbox: Mutex>, - /// Serializes settings mutations (global and sandbox) to prevent - /// read-modify-write races. Held for the duration of any setting - /// set/delete operation, including the precedence check on sandbox - /// mutations that reads global state. - pub settings_mutex: tokio::sync::Mutex<()>, - /// Registry of active supervisor sessions and pending relay channels. /// /// Stored as `Arc` so compiled compute drivers can be constructed before @@ -433,7 +427,6 @@ impl ServerState { telemetry: telemetry::TelemetryState::new(), ssh_connections_by_token: Mutex::new(HashMap::new()), ssh_connections_by_sandbox: Mutex::new(HashMap::new()), - settings_mutex: tokio::sync::Mutex::new(()), supervisor_sessions, gateway_shutting_down: AtomicBool::new(false), replica_id, diff --git a/crates/openshell-server/src/persistence/mod.rs b/crates/openshell-server/src/persistence/mod.rs index 1721cf7931..5e929b0044 100644 --- a/crates/openshell-server/src/persistence/mod.rs +++ b/crates/openshell-server/src/persistence/mod.rs @@ -4,6 +4,7 @@ //! Persistence layer for `OpenShell` Server. mod legacy_time_wire; +pub mod mutation_lock; mod postgres; mod sqlite; @@ -17,6 +18,7 @@ use rand::Rng; use std::collections::HashMap; use thiserror::Error; +pub use mutation_lock::{LockMode, MutationLockKey, MutationLockSet}; pub use postgres::PostgresStore; pub use sqlite::SqliteStore; @@ -58,6 +60,10 @@ pub enum PersistenceError { Conflict { current_resource_version: Option, }, + /// The mutation lock was not acquired before its deadline, so this + /// request's guarded writes did not run; the operation is safe to retry. + #[error("mutation lock timeout: {0}")] + LockTimeout(String), } impl PersistenceError { @@ -202,6 +208,26 @@ pub struct DistributedMutationGuard { _postgres: Option, } +/// RAII guard for the database-backed SSH identity lock. +pub struct SshIdentityMutationGuard { + _postgres: Option, +} + +#[cfg(test)] +impl DistributedMutationGuard { + /// Backend process id of the `PostgreSQL` session holding the locks, or + /// `None` on `SQLite`. + pub(crate) async fn postgres_backend_pid(&mut self) -> Option { + let Self { + _postgres: postgres, + } = self; + match postgres { + Some(guard) => Some(guard.backend_pid().await), + None => None, + } + } +} + /// Trait for inferring an object type string from a message type. pub trait ObjectType { fn object_type() -> &'static str; @@ -285,30 +311,41 @@ impl Store { /// Serialize mutations whose invariants span multiple persisted objects. /// /// `SQLite` deployments are single-replica and use only the caller's local - /// mutex. `PostgreSQL` deployments additionally hold a session-level - /// advisory lock so concurrent gateway replicas cannot validate and write - /// the same cross-object invariant independently. + /// locks. `PostgreSQL` deployments additionally hold `locks` as + /// session-level advisory locks, taken in ascending key order on one + /// connection from the dedicated lock pool, so concurrent gateway replicas + /// cannot validate and write the same cross-object invariant + /// independently. Fails with [`PersistenceError::LockTimeout`] when the + /// locks are not acquired by `deadline`, and with + /// [`PersistenceError::Database`] when `PostgreSQL` does not open a lock + /// connection in at least [`mutation_lock::LOCK_CONNECTION_MIN_BUDGET`]. pub async fn acquire_distributed_mutation_guard( &self, + locks: &MutationLockSet, + deadline: tokio::time::Instant, ) -> PersistenceResult { match self { Self::Postgres(store) => Ok(DistributedMutationGuard { - _postgres: Some(store.acquire_cross_object_lock().await?), + _postgres: Some(store.acquire_mutation_locks(locks, deadline).await?), }), Self::Sqlite(_) => Ok(DistributedMutationGuard { _postgres: None }), } } - /// Independent of the cross-object lock: creation already holds that - /// lock when it provisions a supervisor's durable SSH identity. + /// Independent of the mutation locks: creation already holds its sandbox + /// mutation guard when it provisions a supervisor's durable SSH identity. pub(crate) async fn acquire_ssh_identity_mutation_guard( &self, - ) -> PersistenceResult { + ) -> PersistenceResult { match self { - Self::Postgres(store) => Ok(DistributedMutationGuard { - _postgres: Some(store.acquire_mutation_lock(0x4f53_5348_484f_5354).await?), + Self::Postgres(store) => Ok(SshIdentityMutationGuard { + _postgres: Some( + store + .acquire_data_pool_lock(mutation_lock::SSH_IDENTITY_LOCK_KEY) + .await?, + ), }), - Self::Sqlite(_) => Ok(DistributedMutationGuard { _postgres: None }), + Self::Sqlite(_) => Ok(SshIdentityMutationGuard { _postgres: None }), } } @@ -1041,20 +1078,6 @@ impl Store { .collect() } - /// List and decode protobuf messages across all workspaces, hydrating - /// `resource_version` from the authoritative DB row. - pub async fn list_all_messages( - &self, - limit: u32, - offset: u32, - ) -> PersistenceResult> { - self.list_by_type(T::object_type(), limit, offset) - .await? - .into_iter() - .map(decode_record) - .collect() - } - /// List and decode objects that have a related membership record, with /// pagination. See [`Store::list_with_membership`] for details. pub async fn list_messages_with_membership< @@ -1410,5 +1433,11 @@ pub async fn test_store() -> Store { .expect("in-memory SQLite store should connect") } +#[cfg(test)] +pub mod test_postgres; + +#[cfg(test)] +mod mutation_lock_pg_tests; + #[cfg(test)] mod tests; diff --git a/crates/openshell-server/src/persistence/mutation_lock.rs b/crates/openshell-server/src/persistence/mutation_lock.rs new file mode 100644 index 0000000000..fda0e36537 --- /dev/null +++ b/crates/openshell-server/src/persistence/mutation_lock.rs @@ -0,0 +1,291 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Keys, modes, and deadlines of the mutation locks that serialize +//! cross-object mutations across gateway replicas. +//! +//! The locks form a hierarchy of intention locks. Each key is held shared +//! (S) or exclusive (X): +//! +//! | Mutation | Keys | +//! |---|---| +//! | global policy and settings, platform-scope profiles | X(global) | +//! | providers and workspace-scoped profiles | S(global) X(workspace) | +//! | one sandbox, admin or supervisor | S(global) S(workspace) X(sandbox) | +//! | lifecycle, driver watch, reconcile (process-local only) | S(global) X(sandbox) | +//! | provisioning-deadline reconcile (process-local only) | S(global) S(workspace) X(sandbox) | +//! +//! Ordering rules, which make the scheme deadlock-free: +//! +//! 1. A per-sandbox lifecycle gate, where a path uses one, comes first. +//! 2. Process-local keys follow in ascending `i64` order. +//! 3. On `PostgreSQL`, the same keys follow as session-level advisory locks in +//! ascending order, all on one lock-pool connection. +//! 4. A task never acquires a mutation guard or a local lifecycle lock while +//! it holds one: no nesting and no upgrade. +//! 5. The SSH identity key, held on a data-pool connection, is a leaf: its +//! holders take no mutation guard, local key, or lifecycle gate, so sandbox +//! creation may wait for it while holding its guard. +//! +//! Within each layer every waiter on a key holds only smaller keys, and the +//! local phase ends before the `PostgreSQL` phase starts, so no wait-for cycle +//! can form. +//! +//! The global key is the legacy cross-object key. Gateways from earlier +//! releases hold it exclusively for every mutation, which conflicts with every +//! scope of this release, so mixed-version fleets stay mutually exclusive +//! during a rolling upgrade. + +use sha2::{Digest, Sha256}; +use std::collections::BTreeMap; +use std::time::Duration; + +/// Advisory-lock key of the global mutation lock. +/// +/// Never change this value: gateways from earlier releases hold it +/// exclusively for every cross-object mutation, and a rolling upgrade relies +/// on old and new replicas excluding each other through it. The bytes spell +/// "OPENSHLL" and stay within `PostgreSQL`'s signed 64-bit key space. +pub const GLOBAL_MUTATION_LOCK_KEY: i64 = 0x4f50_454e_5348_4c4c; + +/// Advisory-lock key that serializes sandbox SSH host identity provisioning +/// and cleanup across replicas. +/// +/// Never change this value: every gateway in a fleet must take the same key. +/// It is held on a data-pool connection, not the lock pool, because sandbox +/// creation takes it while holding its mutation guard. The bytes spell +/// "OSSHHOST". +pub const SSH_IDENTITY_LOCK_KEY: i64 = 0x4f53_5348_484f_5354; + +/// Upper bound on acquiring one mutation lock set. +/// +/// Holders validate and write, and some guarded sections also call the +/// compute driver, credential driver, middleware, or profile sources. A wait +/// this long therefore means a stuck replica, a slow dependency, an +/// overloaded database, or a lock pool exhausted by such holders on this +/// replica; failing beats blocking mutations indefinitely. Keep +/// [`MUTATION_LOCK_TIMEOUT_SETTING`] in sync. +pub const MUTATION_LOCK_TIMEOUT: Duration = Duration::from_secs(10); + +/// [`MUTATION_LOCK_TIMEOUT`] as a `PostgreSQL` `lock_timeout` value. +pub const MUTATION_LOCK_TIMEOUT_SETTING: &str = "10s"; + +/// Opening a lock connection can take this long, so a timeout with less time left is contention. +pub const LOCK_CONNECTION_MIN_BUDGET: Duration = Duration::from_secs(1); + +/// Size of the dedicated `PostgreSQL` lock pool. +/// +/// Lock connections come from their own pool so that guard holders can never +/// starve the data pool their critical sections need. Each replica opens at +/// most 10 data plus 4 lock connections. A cancelled acquisition frees its +/// slot at once, but its backend can stay until its `lock_timeout` while the +/// pool opens a replacement, so size `max_connections` with headroom for +/// rollouts as the high-availability guide describes +/// (`(2 × replicas + surge) × 14`). Each guard holds +/// one lock connection, so a replica sustains about 4 / c guarded operations +/// per second, where c is how long one guard is held. +pub(super) const MUTATION_LOCK_POOL_MAX_CONNECTIONS: u32 = 4; + +/// Domain separator hashed into every derived key. +const KEY_DOMAIN: &[u8] = b"openshell/mutation-lock/v1"; + +/// Advisory-lock key of the one-time time-payload migration +/// (`PostgresStore::migrate_legacy_time_payloads`). Derived keys never use it. +const TIME_PAYLOAD_MIGRATION_LOCK_KEY: i64 = 3052; + +/// Mode in which a mutation lock key is held. `Shared` sorts first, so the +/// maximum of two modes is the stronger one. +#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)] +pub enum LockMode { + Shared, + Exclusive, +} + +/// A mutation lock key. +#[derive(Clone, Copy, Debug)] +pub enum MutationLockKey<'a> { + /// The fleet-wide key, [`GLOBAL_MUTATION_LOCK_KEY`]. + Global, + /// One workspace, by name. + Workspace(&'a str), + /// One sandbox, by stable id. + Sandbox(&'a str), +} + +impl MutationLockKey<'_> { + /// The `PostgreSQL` advisory-lock key, also used by the process-local lock + /// table. + /// + /// Derived keys are the first 8 bytes, as a big-endian `i64`, of + /// `SHA-256(KEY_DOMAIN || 0 || kind || 0 || value)`. They are computed in + /// Rust so every replica and every `PostgreSQL` version agrees on them. A + /// hash collision only over-serializes. + pub fn advisory_key(self) -> i64 { + match self { + Self::Global => GLOBAL_MUTATION_LOCK_KEY, + Self::Workspace(workspace) => derived_key(b"workspace", workspace), + Self::Sandbox(sandbox_id) => derived_key(b"sandbox", sandbox_id), + } + } +} + +fn derived_key(kind: &[u8], value: &str) -> i64 { + let digest = Sha256::new() + .chain_update(KEY_DOMAIN) + .chain_update([0]) + .chain_update(kind) + .chain_update([0]) + .chain_update(value.as_bytes()) + .finalize(); + let mut prefix = [0_u8; 8]; + prefix.copy_from_slice(&digest[..8]); + avoid_reserved(i64::from_be_bytes(prefix)) +} + +/// Keep derived keys off the global, SSH identity, and migration keys. +const fn avoid_reserved(key: i64) -> i64 { + if key == GLOBAL_MUTATION_LOCK_KEY + || key == SSH_IDENTITY_LOCK_KEY + || key == TIME_PAYLOAD_MIGRATION_LOCK_KEY + { + key ^ 1 + } else { + key + } +} + +/// The keys one mutation holds, each in its strongest requested mode. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct MutationLockSet { + entries: BTreeMap, +} + +impl MutationLockSet { + /// Process-local lock set of a lifecycle, driver-watch, or reconcile path: + /// S(global) X(sandbox). + pub fn sandbox_lifecycle(sandbox_id: &str) -> Self { + let mut set = Self::default(); + set.insert(MutationLockKey::Global, LockMode::Shared); + set.insert(MutationLockKey::Sandbox(sandbox_id), LockMode::Exclusive); + set + } + + pub fn insert(&mut self, key: MutationLockKey<'_>, mode: LockMode) { + self.insert_raw(key.advisory_key(), mode); + } + + fn insert_raw(&mut self, key: i64, mode: LockMode) { + self.entries + .entry(key) + .and_modify(|held| *held = (*held).max(mode)) + .or_insert(mode); + } + + /// Keys in ascending order, the only acquisition order. + pub fn iter(&self) -> impl Iterator + '_ { + self.entries.iter().map(|(key, mode)| (*key, *mode)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn global_key_is_the_legacy_cross_object_key() { + assert_eq!( + MutationLockKey::Global.advisory_key(), + 0x4f50_454e_5348_4c4c + ); + } + + #[test] + fn timeout_setting_matches_duration() { + assert_eq!( + format!("{}s", MUTATION_LOCK_TIMEOUT.as_secs()), + MUTATION_LOCK_TIMEOUT_SETTING + ); + } + + #[test] + fn lock_set_iterates_in_ascending_key_order() { + let mut set = MutationLockSet::default(); + set.insert_raw(5, LockMode::Exclusive); + set.insert_raw(-3, LockMode::Exclusive); + + assert_eq!( + set.iter().collect::>(), + vec![(-3, LockMode::Exclusive), (5, LockMode::Exclusive)] + ); + } + + #[test] + fn derived_keys_match_golden_values() { + // Computed independently from the documented byte layout. Changing + // any of them breaks mutual exclusion with running replicas. + assert_eq!( + MutationLockKey::Workspace("default").advisory_key(), + 4_171_374_605_116_754_083 + ); + assert_eq!( + MutationLockKey::Workspace("team-a").advisory_key(), + 4_635_337_207_968_654_063 + ); + assert_eq!( + MutationLockKey::Sandbox("00000000-0000-0000-0000-000000000001").advisory_key(), + -542_384_872_970_356_635 + ); + assert_eq!( + MutationLockKey::Sandbox("sb-1").advisory_key(), + -7_385_842_969_463_770_825 + ); + } + + #[test] + fn derived_keys_separate_kinds() { + assert_ne!( + MutationLockKey::Workspace("x").advisory_key(), + MutationLockKey::Sandbox("x").advisory_key() + ); + } + + #[test] + fn reserved_keys_are_remapped() { + assert_ne!( + avoid_reserved(GLOBAL_MUTATION_LOCK_KEY), + GLOBAL_MUTATION_LOCK_KEY + ); + assert_ne!(avoid_reserved(SSH_IDENTITY_LOCK_KEY), SSH_IDENTITY_LOCK_KEY); + assert_eq!(avoid_reserved(TIME_PAYLOAD_MIGRATION_LOCK_KEY), 3053); + assert_eq!(avoid_reserved(42), 42); + } + + #[test] + fn lock_set_keeps_strongest_mode() { + let mut set = MutationLockSet::default(); + set.insert_raw(5, LockMode::Shared); + set.insert_raw(-3, LockMode::Exclusive); + set.insert_raw(5, LockMode::Exclusive); + set.insert_raw(-3, LockMode::Shared); + + assert_eq!( + set.iter().collect::>(), + vec![(-3, LockMode::Exclusive), (5, LockMode::Exclusive)] + ); + } + + #[test] + fn sandbox_lifecycle_set_is_shared_global_exclusive_sandbox() { + let set = MutationLockSet::sandbox_lifecycle("sb-1"); + + let mut expected = vec![ + (GLOBAL_MUTATION_LOCK_KEY, LockMode::Shared), + ( + MutationLockKey::Sandbox("sb-1").advisory_key(), + LockMode::Exclusive, + ), + ]; + expected.sort_unstable(); + assert_eq!(set.iter().collect::>(), expected); + } +} diff --git a/crates/openshell-server/src/persistence/mutation_lock_pg_tests.rs b/crates/openshell-server/src/persistence/mutation_lock_pg_tests.rs new file mode 100644 index 0000000000..68f8476e29 --- /dev/null +++ b/crates/openshell-server/src/persistence/mutation_lock_pg_tests.rs @@ -0,0 +1,932 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! `PostgreSQL` tests of the mutation advisory locks: exclusion between two +//! stores (as between two gateway replicas), exclusion against the legacy +//! global key, cleanup after timed-out and cancelled acquisitions, the +//! lock-pool bound, and lock connections that `PostgreSQL` does not open. +//! +//! Advisory locks are database-wide, not per schema, so every test uses +//! random workspace and sandbox ids, and `mise run test:rust:postgres` runs +//! the tests one at a time. The stores here run no migrations: the schema only +//! scopes their connections. + +use super::mutation_lock::{ + GLOBAL_MUTATION_LOCK_KEY, LOCK_CONNECTION_MIN_BUDGET, MUTATION_LOCK_POOL_MAX_CONNECTIONS, + MUTATION_LOCK_TIMEOUT, MutationLockKey, MutationLockSet, +}; +use super::postgres::LOCK_CONNECTION_RELEASE_TIMEOUT; +use super::test_postgres::TestSchema; +use super::{ + DistributedMutationGuard, PersistenceError, PersistenceResult, PostgresStore, Store, + map_db_error, +}; +use crate::compute::MutationScope; +use sqlx::{Connection, PgConnection, PgPool}; +use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpListener; +use tokio::sync::watch; +use tokio::task::{JoinHandle, JoinSet}; +use tokio::time::Instant; + +/// Deadline of an acquisition that must time out. +const EXPECTED_TIMEOUT: Duration = Duration::from_millis(300); +/// Deadline of an acquisition that must succeed, possibly after a conflicting +/// holder releases. +const PROCEEDS_WITHIN: Duration = Duration::from_secs(5); +/// Deadline of an acquisition that a test interrupts while it waits. +const INTERRUPTED_DEADLINE: Duration = Duration::from_secs(1); +/// How long after an interrupted acquisition's deadline its keys may stay held. +const RELEASED_WITHIN: Duration = Duration::from_secs(1); +/// The client-side backstop fires this long after a lock statement's deadline. +const CLIENT_BACKSTOP_GRACE: Duration = Duration::from_millis(500); +/// How long each holder keeps its locks in the throughput tests. +const HOLD: Duration = Duration::from_millis(20); +const POLL_INTERVAL: Duration = Duration::from_millis(10); + +fn random_id(kind: &str) -> String { + format!("{kind}-{}", uuid::Uuid::new_v4()) +} + +/// Drop client traffic while keeping sockets open, like a failed network path. +/// Client EOF still closes the upstream socket so `PostgreSQL` can release locks. +struct StallingProxy { + url: String, + stalled: watch::Sender, + task: JoinHandle<()>, +} + +impl StallingProxy { + async fn start(database_url: &str) -> Self { + let mut url = url::Url::parse(database_url).unwrap(); + let host = url.host_str().expect("TCP PostgreSQL host").to_owned(); + let port = url.port().unwrap_or(5432); + let upstream_addresses: Vec<_> = tokio::net::lookup_host((host.as_str(), port)) + .await + .unwrap() + .collect(); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + url.set_host(Some("127.0.0.1")).unwrap(); + url.set_port(Some(listener.local_addr().unwrap().port())) + .unwrap(); + let (stalled, receiver) = watch::channel(false); + let task = tokio::spawn(async move { + let mut connections = JoinSet::new(); + loop { + tokio::select! { + accepted = listener.accept() => { + let (client, _) = accepted.unwrap(); + openshell_core::net::set_tcp_nodelay_best_effort(&client); + let addresses = upstream_addresses.clone(); + let stalled = receiver.clone(); + connections.spawn(async move { + let upstream = openshell_core::net::connect_tcp_nodelay_best_effort( + &addresses, + ).await?; + let (mut client_read, mut client_write) = client.into_split(); + let (mut upstream_read, mut upstream_write) = upstream.into_split(); + let requests = async { + let mut buffer = [0_u8; 8192]; + loop { + let read = client_read.read(&mut buffer).await?; + if read == 0 { + return Ok::<_, std::io::Error>(()); + } + if !*stalled.borrow() { + upstream_write.write_all(&buffer[..read]).await?; + } + } + }; + tokio::select! { + result = requests => result, + result = tokio::io::copy(&mut upstream_read, &mut client_write) => { + result.map(|_| ()) + } + } + }); + } + _ = connections.join_next(), if !connections.is_empty() => {} + } + } + }); + Self { + url: url.into(), + stalled, + task, + } + } +} + +impl Drop for StallingProxy { + fn drop(&mut self) { + self.task.abort(); + } +} + +/// A random sandbox id whose key sorts after the global key and the +/// workspace key. Keys are taken in ascending order, so an acquisition of its +/// sandbox scope already holds S(global) and S(workspace) while it waits for +/// the sandbox key. +fn sandbox_id_locked_last(workspace: &str) -> String { + let taken_first = + GLOBAL_MUTATION_LOCK_KEY.max(MutationLockKey::Workspace(workspace).advisory_key()); + loop { + let sandbox = random_id("sb"); + if MutationLockKey::Sandbox(&sandbox).advisory_key() > taken_first { + return sandbox; + } + } +} + +/// A disposable schema plus an observer pool for `pg_locks`. +struct LockFixture { + schema: TestSchema, + observer: PgPool, +} + +impl LockFixture { + async fn new() -> Self { + let schema = TestSchema::create("lock").await; + let observer = PgPool::connect(schema.url()) + .await + .expect("connect the pg_locks observer"); + Self { schema, observer } + } + + /// A store with its own data and lock pools, like one gateway replica. + /// + /// Its lock pool starts with one idle, connected session, so the first + /// acquisition spends its deadline on locks rather than on connecting. + /// Warming takes S(global), so create stores before any test holder + /// locks. + async fn store(&self, lock_pool_size: u32) -> Store { + let store = Store::Postgres( + PostgresStore::connect_with_lock_pool_size(self.schema.url(), lock_pool_size) + .await + .expect("connect a lock store"), + ); + drop( + acquire_proceeds( + &store, + MutationScope::sandbox(&random_id("ws"), &random_id("sb")), + "warm the lock pool", + ) + .await, + ); + wait_for_idle_lock_connection(&store).await; + store + } + + /// A plain session, like a gateway from an earlier release or a test + /// holder. + async fn raw_session(&self) -> PgConnection { + PgConnection::connect(self.schema.url()) + .await + .expect("connect a raw session") + } + + /// Granted and waiting holders of one bigint advisory key in this + /// database. + async fn lock_count(&self, key: i64) -> i64 { + let (high, low) = key_halves(key); + sqlx::query_scalar( + "SELECT count(*) FROM pg_locks \ + WHERE locktype = 'advisory' AND objsubid = 1 \ + AND classid = $1::bigint::oid AND objid = $2::bigint::oid \ + AND database = (SELECT oid FROM pg_database WHERE datname = current_database())", + ) + .bind(high) + .bind(low) + .fetch_one(&self.observer) + .await + .expect("count advisory locks") + } + + /// Advisory locks held or awaited by one backend. + async fn session_lock_count(&self, pid: i32) -> i64 { + sqlx::query_scalar("SELECT count(*) FROM pg_locks WHERE locktype = 'advisory' AND pid = $1") + .bind(pid) + .fetch_one(&self.observer) + .await + .expect("count a session's advisory locks") + } + + async fn wait_for_lock_count(&self, key: i64, expected: i64, until: Instant, what: &str) { + loop { + let count = self.lock_count(key).await; + if count == expected { + return; + } + assert!( + Instant::now() < until, + "{what}: {count} holders remain, expected {expected}" + ); + tokio::time::sleep(POLL_INTERVAL).await; + } + } + + /// Backend that waits for `key`, once an acquisition blocks on it. + async fn wait_for_waiting_backend(&self, key: i64, until: Instant) -> i32 { + let (high, low) = key_halves(key); + loop { + let waiting: Option = sqlx::query_scalar( + "SELECT pid FROM pg_locks \ + WHERE locktype = 'advisory' AND NOT granted AND objsubid = 1 \ + AND classid = $1::bigint::oid AND objid = $2::bigint::oid \ + AND database = (SELECT oid FROM pg_database WHERE datname = current_database())", + ) + .bind(high) + .bind(low) + .fetch_optional(&self.observer) + .await + .expect("find the waiting backend"); + if let Some(pid) = waiting { + return pid; + } + assert!( + Instant::now() < until, + "no backend waited for the sandbox key before the deadline" + ); + tokio::time::sleep(POLL_INTERVAL).await; + } + } + + /// Spawn an acquisition of `Sandbox { workspace, sandbox }` on `store`, + /// where another session holds the sandbox key, and wait until its + /// backend holds S(global) and S(workspace) and waits for the sandbox + /// key. Returns the task and the waiting backend's pid. + async fn spawn_waiting_acquisition( + &self, + store: &Store, + workspace: &str, + sandbox: &str, + deadline: Instant, + ) -> (JoinHandle>, i32) { + let acquisition = { + let store = store.clone(); + let set = MutationScope::sandbox(workspace, sandbox).lock_set(); + tokio::spawn(async move { + store + .acquire_distributed_mutation_guard(&set, deadline) + .await + .map(drop) + }) + }; + let pid = self + .wait_for_waiting_backend(MutationLockKey::Sandbox(sandbox).advisory_key(), deadline) + .await; + assert_eq!( + self.lock_count(MutationLockKey::Workspace(workspace).advisory_key()) + .await, + 1, + "the waiting acquisition holds the workspace key" + ); + assert_eq!( + self.lock_count(GLOBAL_MUTATION_LOCK_KEY).await, + 1, + "the waiting acquisition holds the global key" + ); + (acquisition, pid) + } + + async fn finish(self, stores: Vec) { + for store in stores { + store.close().await; + } + self.observer.close().await; + self.schema.drop_schema().await; + } +} + +/// `pg_locks` shows a bigint advisory key as its high and low 32 bits. +fn key_halves(key: i64) -> (i64, i64) { + let [b0, b1, b2, b3, b4, b5, b6, b7] = key.to_be_bytes(); + ( + i64::from(u32::from_be_bytes([b0, b1, b2, b3])), + i64::from(u32::from_be_bytes([b4, b5, b6, b7])), + ) +} + +async fn acquire( + store: &Store, + scope: MutationScope<'_>, + within: Duration, +) -> PersistenceResult { + store + .acquire_distributed_mutation_guard(&scope.lock_set(), Instant::now() + within) + .await +} + +async fn acquire_proceeds( + store: &Store, + scope: MutationScope<'_>, + what: &str, +) -> DistributedMutationGuard { + acquire(store, scope, PROCEEDS_WITHIN) + .await + .unwrap_or_else(|error| panic!("{what}: expected the locks, got {error:?}")) +} + +/// Wait until `store`'s lock pool has an idle, connected session. A released +/// lock connection returns to the pool from a background task, and an +/// acquisition that starts before then opens a new connection within its own +/// deadline. +async fn wait_for_idle_lock_connection(store: &Store) { + let Store::Postgres(postgres) = store else { + panic!("the lock tests use PostgreSQL stores"); + }; + let until = Instant::now() + PROCEEDS_WITHIN; + while postgres.lock_pool_idle() == 0 { + assert!( + Instant::now() < until, + "no lock connection returned to the pool" + ); + tokio::time::sleep(POLL_INTERVAL).await; + } +} + +/// Acquire `scope` on an idle lock connection with a deadline that +/// `PostgreSQL` must end with its lock timeout (55P03, "canceling statement +/// due to lock timeout"). A client-side timeout fails the test: a connect +/// that misses the deadline never reaches the advisory locks. +async fn acquire_times_out(store: &Store, scope: MutationScope<'_>, what: &str) { + wait_for_idle_lock_connection(store).await; + match acquire(store, scope, EXPECTED_TIMEOUT).await { + Err(PersistenceError::LockTimeout(detail)) if detail.contains("lock timeout") => {} + Err(error) => panic!("{what}: expected PostgreSQL's lock timeout, got {error:?}"), + Ok(_) => panic!("{what}: expected a lock timeout, but the locks were acquired"), + } +} + +async fn raw_lock(session: &mut PgConnection, key: i64) -> sqlx::Result<()> { + sqlx::query("SELECT pg_advisory_lock($1)") + .bind(key) + .execute(session) + .await + .map(drop) +} + +async fn raw_unlock(session: &mut PgConnection, key: i64) { + let released: bool = sqlx::query_scalar("SELECT pg_advisory_unlock($1)") + .bind(key) + .fetch_one(session) + .await + .expect("unlock the raw session's key"); + assert!(released, "the raw session held the key"); +} + +/// Spawn one holder per lock set, spread over `stores`, each keeping its +/// locks for [`HOLD`]; returns how long all of them took. +async fn hold_concurrently(stores: &[Store], sets: Vec) -> Duration { + let started = Instant::now(); + let holders: Vec<_> = sets + .into_iter() + .zip(stores.iter().cycle()) + .map(|(set, store)| { + let store = store.clone(); + tokio::spawn(async move { + let guard = store + .acquire_distributed_mutation_guard( + &set, + Instant::now() + MUTATION_LOCK_TIMEOUT, + ) + .await?; + tokio::time::sleep(HOLD).await; + drop(guard); + Ok::<_, PersistenceError>(()) + }) + }) + .collect(); + for holder in holders { + holder + .await + .expect("holder task") + .expect("holder acquires its locks"); + } + started.elapsed() +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_disjoint_sandbox_scopes_hold_concurrently_across_stores() { + let fixture = LockFixture::new().await; + let store_a = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let store_b = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let workspace = random_id("ws"); + let (sandbox_1, sandbox_2) = (random_id("sb"), random_id("sb")); + + let held_a = acquire_proceeds( + &store_a, + MutationScope::sandbox(&workspace, &sandbox_1), + "store A", + ) + .await; + let held_b = acquire_proceeds( + &store_b, + MutationScope::sandbox(&workspace, &sandbox_2), + "store B, another sandbox in the same workspace", + ) + .await; + assert_eq!( + fixture + .lock_count(MutationLockKey::Workspace(&workspace).advisory_key()) + .await, + 2, + "both sessions hold the workspace key shared at once" + ); + + drop(held_a); + drop(held_b); + fixture.finish(vec![store_a, store_b]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_stalled_release_closes_session_and_recovers_pool() { + let fixture = LockFixture::new().await; + let proxy = StallingProxy::start(fixture.schema.url()).await; + let store = Store::Postgres( + PostgresStore::connect_with_lock_pool_size(&proxy.url, 1) + .await + .unwrap(), + ); + let workspace = random_id("ws"); + let sandbox = random_id("sb"); + let scope = MutationScope::sandbox(&workspace, &sandbox); + let held = acquire_proceeds(&store, scope, "before the network stalls").await; + let key = MutationLockKey::Sandbox(&sandbox).advisory_key(); + assert_eq!(fixture.lock_count(key).await, 1); + + proxy.stalled.send_replace(true); + drop(held); + + // The old pool return waits forever for pg_advisory_unlock_all. Dropping + // a guard must instead bound cleanup and close the stalled connection. + fixture + .wait_for_lock_count( + key, + 0, + Instant::now() + LOCK_CONNECTION_RELEASE_TIMEOUT + Duration::from_secs(2), + "a stalled release must close its lock session", + ) + .await; + let Store::Postgres(postgres) = &store else { + unreachable!(); + }; + assert_eq!(postgres.lock_pool_size(), 0, "the pool permit is recovered"); + + proxy.stalled.send_replace(false); + drop(acquire_proceeds(&store, scope, "after network recovery").await); + fixture.finish(vec![store]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_workspace_exclusive_blocks_same_workspace_sandbox_only() { + let fixture = LockFixture::new().await; + let store_a = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let store_b = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let (workspace_1, workspace_2) = (random_id("ws"), random_id("ws")); + let (sandbox_x, sandbox_y) = (random_id("sb"), random_id("sb")); + + let held = acquire_proceeds(&store_a, MutationScope::Workspace(&workspace_1), "store A").await; + acquire_times_out( + &store_b, + MutationScope::sandbox(&workspace_1, &sandbox_x), + "a sandbox in the held workspace", + ) + .await; + drop( + acquire_proceeds( + &store_b, + MutationScope::sandbox(&workspace_2, &sandbox_y), + "a sandbox in another workspace", + ) + .await, + ); + + drop(held); + drop( + acquire_proceeds( + &store_b, + MutationScope::sandbox(&workspace_1, &sandbox_x), + "the sandbox after the workspace holder releases", + ) + .await, + ); + fixture.finish(vec![store_a, store_b]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_global_exclusive_blocks_every_scope() { + let fixture = LockFixture::new().await; + let store_a = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let store_b = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let workspace = random_id("ws"); + let sandbox = random_id("sb"); + let scopes = [ + MutationScope::Global, + MutationScope::Workspace(""), + MutationScope::Workspace(&workspace), + MutationScope::sandbox(&workspace, &sandbox), + ]; + + let held = acquire_proceeds(&store_a, MutationScope::Global, "store A").await; + for scope in scopes { + acquire_times_out(&store_b, scope, &format!("{scope:?} behind the global key")).await; + } + + drop(held); + for scope in scopes { + drop(acquire_proceeds(&store_b, scope, &format!("{scope:?} after release")).await); + } + fixture.finish(vec![store_a, store_b]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_legacy_global_key_holder_excludes_new_scopes() { + let fixture = LockFixture::new().await; + let store = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let workspace = random_id("ws"); + let sandbox = random_id("sb"); + let scope = MutationScope::sandbox(&workspace, &sandbox); + + // A gateway from an earlier release holds the legacy key exclusively. + let mut legacy = fixture.raw_session().await; + raw_lock(&mut legacy, GLOBAL_MUTATION_LOCK_KEY) + .await + .expect("the legacy holder takes the global key"); + acquire_times_out(&store, scope, "a sandbox scope behind a legacy holder").await; + + raw_unlock(&mut legacy, GLOBAL_MUTATION_LOCK_KEY).await; + drop( + acquire_proceeds( + &store, + scope, + "a sandbox scope after the legacy holder releases", + ) + .await, + ); + + legacy.close().await.expect("close the legacy session"); + fixture.finish(vec![store]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_new_scope_holder_excludes_legacy_global_key() { + let fixture = LockFixture::new().await; + let store = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let workspace = random_id("ws"); + let sandbox = random_id("sb"); + + let held = acquire_proceeds( + &store, + MutationScope::sandbox(&workspace, &sandbox), + "the new-release holder", + ) + .await; + let mut legacy = fixture.raw_session().await; + sqlx::query("SET lock_timeout = '200ms'") + .execute(&mut legacy) + .await + .expect("bound the legacy lock wait"); + let Err(error) = raw_lock(&mut legacy, GLOBAL_MUTATION_LOCK_KEY).await else { + panic!("the legacy global lock should wait behind a sandbox scope"); + }; + assert_eq!( + error + .as_database_error() + .and_then(sqlx::error::DatabaseError::code) + .as_deref(), + Some("55P03"), + "{error}" + ); + // Only a guard acquisition turns 55P03 into a lock timeout; the + // generic mapping keeps any other session's 55P03 a database error. + assert!( + matches!(map_db_error(&error), PersistenceError::Database(_)), + "a 55P03 outside a guard acquisition is a database error" + ); + + drop(held); + sqlx::query("SET lock_timeout = '5s'") + .execute(&mut legacy) + .await + .expect("bound the legacy lock wait"); + raw_lock(&mut legacy, GLOBAL_MUTATION_LOCK_KEY) + .await + .expect("the legacy lock after the new-release holder releases"); + raw_unlock(&mut legacy, GLOBAL_MUTATION_LOCK_KEY).await; + + legacy.close().await.expect("close the legacy session"); + fixture.finish(vec![store]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_timeout_releases_held_keys_before_the_holder_does() { + let fixture = LockFixture::new().await; + // One lock connection, so a reused session keeps its backend pid. + let store = fixture.store(1).await; + let workspace = random_id("ws"); + let sandbox = sandbox_id_locked_last(&workspace); + let workspace_key = MutationLockKey::Workspace(&workspace).advisory_key(); + let sandbox_key = MutationLockKey::Sandbox(&sandbox).advisory_key(); + + // The raw session holds only the sandbox key, so the acquisition's own + // backend is the only holder of the global and workspace keys. + let mut holder = fixture.raw_session().await; + raw_lock(&mut holder, sandbox_key) + .await + .expect("the raw session takes the sandbox key"); + + let deadline = Instant::now() + INTERRUPTED_DEADLINE; + let (acquisition, waiting_pid) = fixture + .spawn_waiting_acquisition(&store, &workspace, &sandbox, deadline) + .await; + + let result = tokio::time::timeout_at( + deadline + CLIENT_BACKSTOP_GRACE + RELEASED_WITHIN, + acquisition, + ) + .await + .expect("the acquisition returns by its deadline") + .expect("acquisition task"); + assert!( + matches!(result, Err(PersistenceError::LockTimeout(_))), + "{result:?}" + ); + + // The raw session still holds the sandbox key, yet the global and + // workspace keys the acquisition took are already released. + fixture + .wait_for_lock_count( + workspace_key, + 0, + deadline + RELEASED_WITHIN, + "workspace key after the timeout", + ) + .await; + fixture + .wait_for_lock_count( + GLOBAL_MUTATION_LOCK_KEY, + 0, + deadline + RELEASED_WITHIN, + "global key after the timeout", + ) + .await; + assert_eq!(fixture.lock_count(sandbox_key).await, 1); + + // A server-side timeout returns the healthy session to the pool instead + // of closing it. + let mut reused = acquire_proceeds( + &store, + MutationScope::sandbox(&random_id("ws"), &random_id("sb")), + "a disjoint scope on the same lock connection", + ) + .await; + assert_eq!(reused.postgres_backend_pid().await, Some(waiting_pid)); + drop(reused); + + raw_unlock(&mut holder, sandbox_key).await; + holder.close().await.expect("close the raw session"); + fixture.finish(vec![store]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_cancelled_acquisition_releases_held_keys_by_its_deadline() { + let fixture = LockFixture::new().await; + let store = fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await; + let workspace = random_id("ws"); + let sandbox = sandbox_id_locked_last(&workspace); + let workspace_key = MutationLockKey::Workspace(&workspace).advisory_key(); + let sandbox_key = MutationLockKey::Sandbox(&sandbox).advisory_key(); + + let mut holder = fixture.raw_session().await; + raw_lock(&mut holder, sandbox_key) + .await + .expect("the raw session takes the sandbox key"); + + // Cancel mid-wait, while the backend holds the global and workspace keys. + let deadline = Instant::now() + INTERRUPTED_DEADLINE; + let (acquisition, _) = fixture + .spawn_waiting_acquisition(&store, &workspace, &sandbox, deadline) + .await; + acquisition.abort(); + let Err(error) = acquisition.await else { + panic!("the cancelled acquisition should not finish"); + }; + assert!(error.is_cancelled()); + + // The closed session's backend does not notice the disconnect while it + // waits, but the statement's own lock_timeout ends the wait at the + // original deadline. Without it the keys would stay held for the full + // 10 s session backstop. + fixture + .wait_for_lock_count( + workspace_key, + 0, + deadline + RELEASED_WITHIN, + "workspace key after the cancellation", + ) + .await; + fixture + .wait_for_lock_count( + GLOBAL_MUTATION_LOCK_KEY, + 0, + deadline + RELEASED_WITHIN, + "global key after the cancellation", + ) + .await; + + raw_unlock(&mut holder, sandbox_key).await; + fixture + .wait_for_lock_count( + sandbox_key, + 0, + Instant::now() + PROCEEDS_WITHIN, + "sandbox key after the raw session unlocks", + ) + .await; + holder.close().await.expect("close the raw session"); + fixture.finish(vec![store]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_released_connection_is_reused_without_residual_locks() { + let fixture = LockFixture::new().await; + let store = fixture.store(1).await; + let workspace = random_id("ws"); + let sandbox = random_id("sb"); + let scope = MutationScope::sandbox(&workspace, &sandbox); + + let mut first = acquire_proceeds(&store, scope, "the first acquisition").await; + let pid = first + .postgres_backend_pid() + .await + .expect("a PostgreSQL guard"); + assert_eq!(fixture.session_lock_count(pid).await, 3); + drop(first); + + // Returning the connection unlocks everything it held. + let until = Instant::now() + PROCEEDS_WITHIN; + while fixture.session_lock_count(pid).await != 0 { + assert!( + Instant::now() < until, + "the released session still holds advisory locks" + ); + tokio::time::sleep(POLL_INTERVAL).await; + } + + let mut second = acquire_proceeds(&store, scope, "the second acquisition").await; + assert_eq!(second.postgres_backend_pid().await, Some(pid)); + drop(second); + fixture.finish(vec![store]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_pool_never_exceeds_configured_size() { + let fixture = LockFixture::new().await; + let postgres = PostgresStore::connect(fixture.schema.url()) + .await + .expect("connect a store with the production lock pool"); + let store = Store::Postgres(postgres.clone()); + let workspace = random_id("ws"); + + let holders: Vec<_> = (0..32) + .map(|_| { + let store = store.clone(); + let workspace = workspace.clone(); + tokio::spawn(async move { + let sandbox = random_id("sb"); + let guard = acquire( + &store, + MutationScope::sandbox(&workspace, &sandbox), + PROCEEDS_WITHIN, + ) + .await?; + tokio::time::sleep(HOLD).await; + drop(guard); + Ok::<_, PersistenceError>(()) + }) + }) + .collect(); + let mut largest = 0; + while !holders.iter().all(JoinHandle::is_finished) { + largest = largest.max(postgres.lock_pool_size()); + tokio::time::sleep(Duration::from_millis(1)).await; + } + for holder in holders { + holder + .await + .expect("holder task") + .expect("every holder acquires its locks"); + } + largest = largest.max(postgres.lock_pool_size()); + + assert_eq!( + largest, MUTATION_LOCK_POOL_MAX_CONNECTIONS, + "32 concurrent holders fill the lock pool without exceeding it" + ); + fixture.finish(vec![store]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_connection_failure_is_not_a_lock_timeout() { + let fixture = LockFixture::new().await; + + // Every lock connection is checked out, so the acquisition waits for one + // to come back: lock contention. The deadline leaves time to open a + // connection, so only the full pool makes this a lock timeout. + let saturated = fixture.store(1).await; + let held = acquire_proceeds( + &saturated, + MutationScope::sandbox(&random_id("ws"), &random_id("sb")), + "the only lock connection", + ) + .await; + match acquire( + &saturated, + MutationScope::sandbox(&random_id("ws"), &random_id("sb")), + LOCK_CONNECTION_MIN_BUDGET + EXPECTED_TIMEOUT, + ) + .await + { + Err(PersistenceError::LockTimeout(detail)) => { + assert_eq!(detail, "waiting for a mutation lock connection"); + } + Err(error) => panic!("expected a lock timeout from a full pool, got {error:?}"), + Ok(_) => panic!("a full lock pool handed out a second connection"), + } + drop(held); + + // The pool has room and the deadline leaves time to open a connection, + // but PostgreSQL never answers one: a database failure, not contention. + let proxy = StallingProxy::start(fixture.schema.url()).await; + let unanswered = Store::Postgres( + PostgresStore::connect_with_lock_pool_size(&proxy.url, 1) + .await + .expect("connect a store through the proxy"), + ); + proxy.stalled.send_replace(true); + match acquire( + &unanswered, + MutationScope::sandbox(&random_id("ws"), &random_id("sb")), + LOCK_CONNECTION_MIN_BUDGET + EXPECTED_TIMEOUT, + ) + .await + { + Err(PersistenceError::Database(detail)) => assert!( + detail.starts_with("could not open a mutation lock connection"), + "{detail}" + ), + Err(error) => panic!("expected a database error, got {error:?}"), + Ok(_) => panic!("a stalled PostgreSQL opened a lock connection"), + } + proxy.stalled.send_replace(false); + + fixture.finish(vec![saturated, unanswered]).await; +} + +#[tokio::test] +#[ignore = "requires a disposable PostgreSQL database in OPENSHELL_TEST_POSTGRES_URL; run mise run test:rust:postgres"] +async fn postgres_mutation_lock_unrelated_sandboxes_outpace_global_serialization() { + const HOLDERS: usize = 64; + let fixture = LockFixture::new().await; + let stores = vec![ + fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await, + fixture.store(MUTATION_LOCK_POOL_MAX_CONNECTIONS).await, + ]; + let workspace = random_id("ws"); + + let global = hold_concurrently( + &stores, + (0..HOLDERS) + .map(|_| MutationScope::Global.lock_set()) + .collect(), + ) + .await; + let sandboxes = hold_concurrently( + &stores, + (0..HOLDERS) + .map(|_| MutationScope::sandbox(&workspace, &random_id("sb")).lock_set()) + .collect(), + ) + .await; + + // Global holders serialize (about HOLDERS x HOLD); distinct sandboxes + // share the two lock pools. Compare the runs, not absolute times. + assert!( + sandboxes * 3 < global, + "distinct sandboxes took {sandboxes:?}, global serialization took {global:?}" + ); + fixture.finish(stores).await; +} diff --git a/crates/openshell-server/src/persistence/postgres.rs b/crates/openshell-server/src/persistence/postgres.rs index edaeac2ff5..bfb1084e0a 100644 --- a/crates/openshell-server/src/persistence/postgres.rs +++ b/crates/openshell-server/src/persistence/postgres.rs @@ -1,6 +1,10 @@ // SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 +use super::mutation_lock::{ + LOCK_CONNECTION_MIN_BUDGET, LockMode, MUTATION_LOCK_POOL_MAX_CONNECTIONS, + MUTATION_LOCK_TIMEOUT, MUTATION_LOCK_TIMEOUT_SETTING, MutationLockSet, +}; use super::{ DraftChunkRecord, ObjectCursor, ObjectListQuery, ObjectRecord, PersistenceError, PersistenceResult, PolicyRecord, WriteCondition, WriteResult, current_time_ms, map_db_error, @@ -17,6 +21,9 @@ use prost::Message; use sqlx::pool::PoolConnection; use sqlx::postgres::PgPoolOptions; use sqlx::{Connection, PgPool, Postgres, QueryBuilder, Row}; +use std::sync::Arc; +use std::sync::atomic::{AtomicU32, AtomicU64, Ordering}; +use std::time::Duration; static POSTGRES_MIGRATOR: sqlx::migrate::Migrator = sqlx::migrate!("./migrations/postgres"); @@ -33,34 +40,269 @@ use super::{DELETE_MANY_BATCH_SIZE, DRAFT_CHUNK_OBJECT_TYPE, POLICY_OBJECT_TYPE} #[derive(Debug, Clone)] pub struct PostgresStore { pool: PgPool, + /// Dedicated connections for session-level mutation advisory locks. + lock_pool: PgPool, + /// How acquisitions and guards use the lock pool's connections. + lock_usage: Arc, +} + +/// Lock-pool connections checked out by acquisitions and guards. +#[derive(Debug)] +struct LockPoolUsage { + /// Connections the lock pool may open. + max: u32, + /// Checked-out connections, each counted until its pool permit is + /// released. + in_use: AtomicU32, + /// How many times `in_use` reached `max`. + became_full: AtomicU64, } -// Stable cluster-wide key for serializing sandbox/provider cross-object -// mutations. The bytes spell "OPENSHLL" and stay within PostgreSQL's signed -// 64-bit advisory-lock key space. -const CROSS_OBJECT_ADVISORY_LOCK_KEY: i64 = 0x4f50_454e_5348_4c4c; +/// What a wait for a lock connection saw of the pool when it started. +#[derive(Clone, Copy)] +struct LockPoolSnapshot { + full: bool, + became_full: u64, +} + +impl LockPoolUsage { + fn new(max: u32) -> Self { + Self { + max, + in_use: AtomicU32::new(0), + became_full: AtomicU64::new(0), + } + } + + fn snapshot(&self) -> LockPoolSnapshot { + // Read the counter first: if the pool fills between the two reads, + // either read shows it. + let became_full = self.became_full.load(Ordering::Acquire); + LockPoolSnapshot { + full: self.in_use.load(Ordering::Acquire) >= self.max, + became_full, + } + } -// Bounds the wait for the cross-object lock. The holder only validates and -// writes, so a wait this long means a stuck replica; failing beats blocking -// every sandbox and provider mutation in the fleet indefinitely. -const CROSS_OBJECT_ADVISORY_LOCK_TIMEOUT: &str = "10s"; + /// Whether every connection was checked out at some point since `start`. + /// A connection handed from one guard to the next leaves the count one + /// short only briefly, so a contended pool keeps filling up again. + fn was_full_since(&self, start: LockPoolSnapshot) -> bool { + start.full || self.became_full.load(Ordering::Acquire) != start.became_full + } +} +/// Detail of a lock-pool connection that `PostgreSQL` did not open by the +/// deadline, though it had at least [`LOCK_CONNECTION_MIN_BUDGET`]. `SQLx` +/// retries refused connections, `too_many_connections` (53300), and +/// `cannot_connect_now` (57P03) silently until then. +const LOCK_CONNECTION_NOT_OPENED: &str = "could not open a mutation lock connection: \ + PostgreSQL refused the connection, had no free connection slot, was starting up, \ + or did not answer before the deadline"; + +/// How long the client waits past the caller's deadline for a lock statement. +/// +/// Each statement sets `lock_timeout` to the remaining deadline, so +/// `PostgreSQL` ends the wait on time; this timer only catches a stalled +/// connection. +const LOCK_STATEMENT_CLIENT_GRACE: Duration = Duration::from_millis(500); + +/// Bound the complete pool return, including the unlock hook and `SQLx`'s ping. +/// A stalled connection must not retain a lock-pool permit indefinitely. +pub(super) const LOCK_CONNECTION_RELEASE_TIMEOUT: Duration = Duration::from_secs(5); + +/// Holds the mutation advisory locks of one acquisition. +/// +/// Returned to the lock pool on drop; `after_release` runs +/// `pg_advisory_unlock_all()` before reuse. The entire return is bounded by +/// [`LOCK_CONNECTION_RELEASE_TIMEOUT`]. An acquisition that `PostgreSQL` +/// times out (`55P03`) is returned the same way. A cancelled or otherwise failed +/// acquisition closes its session instead (see [`PendingLockConnection`]). pub(super) struct PostgresAdvisoryLockGuard { - // `close_on_drop` is set before this guard is constructed. Closing the - // dedicated session releases the session-level advisory lock even when a - // request is cancelled or returns early. + connection: PoolConnection, + checkout: LockConnectionCheckout, +} + +/// Holds a session-level advisory lock on a data-pool connection, which is +/// closed on drop. +pub(super) struct PostgresDataPoolLockGuard { _connection: PoolConnection, } +impl Drop for PostgresAdvisoryLockGuard { + fn drop(&mut self) { + return_lock_connection(&mut self.connection, std::mem::take(&mut self.checkout)); + } +} + +/// Counts one checked-out lock-pool connection in [`LockPoolUsage`] until +/// dropped. Its holder drops it only after the connection's pool permit is +/// released. +#[derive(Default)] +struct LockConnectionCheckout(Option>); + +impl LockConnectionCheckout { + fn new(usage: &Arc) -> Self { + if usage.in_use.fetch_add(1, Ordering::AcqRel) + 1 >= usage.max { + usage.became_full.fetch_add(1, Ordering::AcqRel); + } + Self(Some(Arc::clone(usage))) + } +} + +impl Drop for LockConnectionCheckout { + fn drop(&mut self) { + if let Some(usage) = self.0.take() { + usage.in_use.fetch_sub(1, Ordering::AcqRel); + } + } +} + +fn return_lock_connection( + connection: &mut PoolConnection, + checkout: LockConnectionCheckout, +) { + // SQLx transfers the connection and pool permit into this owned future + // immediately. Dropping it on timeout closes the socket and releases the + // permit, including when the unlock hook succeeded but the final ping stalls. + // This relies on sqlx-core 0.9's doc-hidden `PoolConnection::return_to_pool`; + // rerun the `postgres_mutation_lock_stalled_release_*` test on any sqlx bump. + let returning = connection.return_to_pool(); + tokio::spawn(async move { + let _checkout = checkout; + if tokio::time::timeout(LOCK_CONNECTION_RELEASE_TIMEOUT, returning) + .await + .is_err() + { + tracing::warn!( + "timed out returning PostgreSQL mutation lock connection; discarded connection" + ); + } + }); +} + +#[cfg(test)] +impl PostgresAdvisoryLockGuard { + /// Backend process id of the session that holds the locks. + pub(super) async fn backend_pid(&mut self) -> i32 { + let Self { connection, .. } = self; + sqlx::query_scalar("SELECT pg_backend_pid()") + .fetch_one(&mut **connection) + .await + .expect("read the lock session's backend pid") + } +} + +/// A lock-pool connection whose acquisition is still in progress. +/// +/// Dropping it closes the session, which releases every advisory lock the +/// backend holds once its current statement ends. A backend blocked on a lock +/// does not notice the closed socket, so each lock statement carries its own +/// `lock_timeout` to bound that wait by the caller's deadline. +struct PendingLockConnection { + connection: Option>, + checkout: LockConnectionCheckout, +} + +impl PendingLockConnection { + fn new(connection: PoolConnection, usage: &Arc) -> Self { + Self { + connection: Some(connection), + checkout: LockConnectionCheckout::new(usage), + } + } + + fn connection(&mut self) -> &mut PoolConnection { + self.connection + .as_mut() + .expect("pending lock connection is present until disarmed") + } + + /// Keep the session: every lock was acquired. + fn into_guard(mut self) -> PostgresAdvisoryLockGuard { + PostgresAdvisoryLockGuard { + connection: self + .connection + .take() + .expect("pending lock connection is present until disarmed"), + checkout: std::mem::take(&mut self.checkout), + } + } + + /// Return the healthy session to the pool, whose `after_release` unlocks + /// any keys it already holds. + fn release(mut self) { + if let Some(mut connection) = self.connection.take() { + return_lock_connection(&mut connection, std::mem::take(&mut self.checkout)); + } + } +} + +impl Drop for PendingLockConnection { + fn drop(&mut self) { + if let Some(mut connection) = self.connection.take() { + // Still closes the session if the task below never runs. + connection.close_on_drop(); + let checkout = std::mem::take(&mut self.checkout); + // `close` keeps the pool permit until the session is closed, and + // the connection counts as checked out until then. + tokio::spawn(async move { + let _checkout = checkout; + let _ = + tokio::time::timeout(LOCK_CONNECTION_RELEASE_TIMEOUT, connection.close()).await; + }); + } + } +} + impl PostgresStore { pub async fn connect(url: &str) -> PersistenceResult { + Self::connect_with_lock_pool_size(url, MUTATION_LOCK_POOL_MAX_CONNECTIONS).await + } + + pub(super) async fn connect_with_lock_pool_size( + url: &str, + lock_pool_size: u32, + ) -> PersistenceResult { let pool = PgPoolOptions::new() .max_connections(10) .connect(url) .await .map_err(|e| map_db_error(&e))?; + let lock_pool = PgPoolOptions::new() + .max_connections(lock_pool_size) + .min_connections(0) + // Backstop only; callers bound the acquire by their own deadline. + .acquire_timeout(MUTATION_LOCK_TIMEOUT) + .after_connect(|connection, _metadata| { + Box::pin(async move { + // Backstop only: every lock statement sets the remaining + // deadline of its own acquisition. + sqlx::query("SELECT set_config('lock_timeout', $1, false)") + .bind(MUTATION_LOCK_TIMEOUT_SETTING) + .execute(&mut *connection) + .await?; + Ok(()) + }) + }) + .after_release(|connection, _metadata| { + Box::pin(async move { + // Scrub every returned lock connection. On error sqlx closes + // it, and the backend exit releases whatever it held. + sqlx::query("SELECT pg_advisory_unlock_all()") + .execute(&mut *connection) + .await?; + Ok(true) + }) + }) + .connect_lazy(url) + .map_err(|e| map_db_error(&e))?; - Ok(Self { pool }) + Ok(Self { + pool, + lock_pool, + lock_usage: Arc::new(LockPoolUsage::new(lock_pool_size)), + }) } pub async fn migrate(&self) -> PersistenceResult<()> { @@ -114,21 +356,101 @@ impl PostgresStore { conn.ping().await.map_err(|e| map_db_error(&e)) } - pub(super) async fn acquire_cross_object_lock( + /// Acquire `locks` as session-level advisory locks, in ascending key + /// order, on one lock-pool connection. + /// + /// Fails with [`PersistenceError::LockTimeout`] when the lock connections + /// stay checked out, too little time is left to open one, or a lock is + /// not granted by `deadline`, and with [`PersistenceError::Database`] when + /// `PostgreSQL` does not open a lock connection in at least + /// [`LOCK_CONNECTION_MIN_BUDGET`]. + pub(super) async fn acquire_mutation_locks( &self, + locks: &MutationLockSet, + deadline: tokio::time::Instant, ) -> PersistenceResult { - self.acquire_mutation_lock(CROSS_OBJECT_ADVISORY_LOCK_KEY) - .await + let usage_at_start = self.lock_usage.snapshot(); + let budget = deadline.saturating_duration_since(tokio::time::Instant::now()); + let connection = match tokio::time::timeout_at(deadline, self.lock_pool.acquire()).await { + Err(_) | Ok(Err(sqlx::Error::PoolTimedOut)) => { + return Err(self.lock_connection_timeout(usage_at_start, budget)); + } + Ok(Err(error)) => { + return Err(PersistenceError::Database(format!( + "could not open a mutation lock connection: {error}" + ))); + } + Ok(Ok(connection)) => connection, + }; + let mut pending = PendingLockConnection::new(connection, &self.lock_usage); + for (key, mode) in locks.iter() { + let remaining = deadline.saturating_duration_since(tokio::time::Instant::now()); + if remaining < Duration::from_millis(1) { + // Nothing waits server-side; `after_release` unlocks the keys + // taken so far. + pending.release(); + return Err(PersistenceError::LockTimeout( + "waiting for a PostgreSQL advisory lock".into(), + )); + } + // `set_config` and the lock run in one statement. A CTE that calls + // a volatile function is never inlined, and the outer projection + // needs its row, so `lock_timeout` is set before the lock wait + // starts. PostgreSQL then abandons the wait at the caller's + // deadline even if this future is cancelled and the socket closed. + let sql = match mode { + LockMode::Shared => { + "WITH timeout AS (SELECT set_config('lock_timeout', $1, false)) \ + SELECT pg_advisory_lock_shared($2) FROM timeout" + } + LockMode::Exclusive => { + "WITH timeout AS (SELECT set_config('lock_timeout', $1, false)) \ + SELECT pg_advisory_lock($2) FROM timeout" + } + }; + let statement = sqlx::query(sql) + .bind(format!("{}ms", remaining.as_millis())) + .bind(key) + .execute(&mut **pending.connection()); + match tokio::time::timeout_at(deadline + LOCK_STATEMENT_CLIENT_GRACE, statement).await { + Err(_) => { + return Err(PersistenceError::LockTimeout( + "waiting for a PostgreSQL advisory lock".into(), + )); + } + Ok(Err(error)) => { + // 55P03 lock_not_available: the statement's own + // `lock_timeout` ended the wait. The session is healthy: + // return it; `after_release` unlocks the keys already held. + if let Some(db) = error.as_database_error() + && db.code().as_deref() == Some("55P03") + { + let detail = db.message().to_string(); + pending.release(); + return Err(PersistenceError::LockTimeout(detail)); + } + return Err(map_db_error(&error)); + } + Ok(Ok(_)) => {} + } + } + Ok(pending.into_guard()) } - pub(super) async fn acquire_mutation_lock( + /// Acquire `key` as a session-level advisory lock on a data-pool + /// connection instead of the mutation lock pool. Sandbox creation takes + /// the SSH identity lock while it holds its mutation guard, so a second + /// lock-pool connection there could exhaust that pool. The wait uses the + /// mutation lock timeout as its `lock_timeout`. Closing the connection on + /// drop releases the lock. + pub(super) async fn acquire_data_pool_lock( &self, key: i64, - ) -> PersistenceResult { + ) -> PersistenceResult { let mut connection = self.pool.acquire().await.map_err(|e| map_db_error(&e))?; connection.close_on_drop(); sqlx::query("SELECT set_config('lock_timeout', $1, false)") - .bind(CROSS_OBJECT_ADVISORY_LOCK_TIMEOUT) + .bind(MUTATION_LOCK_TIMEOUT_SETTING) .execute(&mut *connection) .await .map_err(|e| map_db_error(&e))?; @@ -137,17 +459,50 @@ impl PostgresStore { .execute(&mut *connection) .await .map_err(|e| map_db_error(&e))?; - Ok(PostgresAdvisoryLockGuard { + Ok(PostgresDataPoolLockGuard { _connection: connection, }) } - /// Test support only: close the underlying connection pool. + /// The error of a lock-pool acquire that ran out of time. + /// + /// If every lock connection was checked out at some point during the + /// wait, the acquire waited for one to come back, which is lock + /// contention. So is a wait that started with less than + /// [`LOCK_CONNECTION_MIN_BUDGET`] left: an earlier wait, such as for + /// local keys, used up the deadline. Otherwise the pool had room the whole + /// time and `PostgreSQL` did not open a new connection. + fn lock_connection_timeout( + &self, + start: LockPoolSnapshot, + budget: Duration, + ) -> PersistenceError { + if budget < LOCK_CONNECTION_MIN_BUDGET || self.lock_usage.was_full_since(start) { + PersistenceError::LockTimeout("waiting for a mutation lock connection".into()) + } else { + PersistenceError::Database(LOCK_CONNECTION_NOT_OPENED.into()) + } + } + + /// Connections the lock pool holds, idle or in use. + #[cfg(test)] + pub(super) fn lock_pool_size(&self) -> u32 { + self.lock_pool.size() + } + + /// Lock-pool connections that are connected and ready for reuse. + #[cfg(test)] + pub(super) fn lock_pool_idle(&self) -> usize { + self.lock_pool.num_idle() + } + + /// Test support only: close the underlying connection pools. /// - /// Do not call from runtime code; this tears down the active pool. + /// Do not call from runtime code; this tears down the active pools. #[cfg(any(test, feature = "test-support"))] pub async fn close(&self) { self.pool.close().await; + self.lock_pool.close().await; } pub async fn put( @@ -1699,3 +2054,138 @@ fn row_to_draft_chunk_record(row: sqlx::postgres::PgRow) -> PersistenceResult Self { + Self { + pool: PgPoolOptions::new() + .max_connections(10) + .connect_lazy(url) + .expect("build the lazy data pool"), + lock_pool: PgPoolOptions::new() + .max_connections(lock_pool_size) + .min_connections(0) + .acquire_timeout(MUTATION_LOCK_TIMEOUT) + .connect_lazy(url) + .expect("build the lazy lock pool"), + lock_usage: Arc::new(LockPoolUsage::new(lock_pool_size)), + } + } + + /// Test support only: a `PostgreSQL` URL on a loopback port where nothing + /// listens, so every connection attempt is refused. + pub(crate) async fn refusing_url_for_tests() -> String { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("reserve a loopback port"); + let port = listener.local_addr().expect("loopback address").port(); + drop(listener); + format!("postgres://openshell@127.0.0.1:{port}/openshell") + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn lock_connection_timeout_is_contention_only_if_the_pool_filled_during_the_wait() { + let store = PostgresStore::connect_lazy_for_tests( + &PostgresStore::refusing_url_for_tests().await, + 2, + ); + let usage = &store.lock_usage; + let is_contention = |error: PersistenceError| match error { + PersistenceError::LockTimeout(detail) => { + assert_eq!(detail, "waiting for a mutation lock connection"); + true + } + PersistenceError::Database(detail) => { + assert_eq!(detail, LOCK_CONNECTION_NOT_OPENED); + false + } + other => panic!("unexpected lock connection error: {other:?}"), + }; + + // The pool never filled: PostgreSQL did not open a connection. + let start = usage.snapshot(); + let first = LockConnectionCheckout::new(usage); + assert!(!is_contention( + store.lock_connection_timeout(start, LOCK_CONNECTION_MIN_BUDGET) + )); + + // The pool was full when the wait started. + let second = LockConnectionCheckout::new(usage); + assert!(is_contention(store.lock_connection_timeout( + usage.snapshot(), + LOCK_CONNECTION_MIN_BUDGET + ))); + + // A connection went back and was checked out again during the wait, + // as when one guard hands its connection to the next. + drop(second); + let start = usage.snapshot(); + drop(LockConnectionCheckout::new(usage)); + assert!(is_contention( + store.lock_connection_timeout(start, LOCK_CONNECTION_MIN_BUDGET) + )); + drop(first); + } + + #[tokio::test] + async fn lock_connection_timeout_without_time_to_connect_is_a_lock_timeout() { + let store = PostgresStore::connect_lazy_for_tests( + &PostgresStore::refusing_url_for_tests().await, + 1, + ); + // The pool never fills, but an earlier wait left too little time to + // open a connection. + let start = store.lock_usage.snapshot(); + let budget = LOCK_CONNECTION_MIN_BUDGET.saturating_sub(Duration::from_millis(1)); + assert!(matches!( + store.lock_connection_timeout(start, budget), + PersistenceError::LockTimeout(_) + )); + + // SQLx retries the refused connection until the deadline. + let deadline = tokio::time::Instant::now() + LOCK_CONNECTION_MIN_BUDGET / 4; + match store + .acquire_mutation_locks(&MutationLockSet::default(), deadline) + .await + { + Err(PersistenceError::LockTimeout(detail)) => { + assert_eq!(detail, "waiting for a mutation lock connection"); + } + Err(error) => panic!("expected a lock timeout, got {error:?}"), + Ok(_) => panic!("nothing listens, yet a lock connection opened"), + } + assert_eq!(store.lock_usage.in_use.load(Ordering::Acquire), 0); + } + + #[tokio::test] + async fn refused_lock_connection_is_a_database_error_not_a_lock_timeout() { + let store = PostgresStore::connect_lazy_for_tests( + &PostgresStore::refusing_url_for_tests().await, + 1, + ); + // SQLx retries a refused connection until the deadline, which leaves + // ample time to open one. + let deadline = + tokio::time::Instant::now() + LOCK_CONNECTION_MIN_BUDGET + Duration::from_millis(300); + match store + .acquire_mutation_locks(&MutationLockSet::default(), deadline) + .await + { + Err(PersistenceError::Database(detail)) => assert!( + detail.starts_with("could not open a mutation lock connection"), + "{detail}" + ), + Err(error) => panic!("expected a database error, got {error:?}"), + Ok(_) => panic!("nothing listens, yet a lock connection opened"), + } + assert_eq!(store.lock_usage.in_use.load(Ordering::Acquire), 0); + } +} diff --git a/crates/openshell-server/src/persistence/test_postgres.rs b/crates/openshell-server/src/persistence/test_postgres.rs new file mode 100644 index 0000000000..8e31d8d80d --- /dev/null +++ b/crates/openshell-server/src/persistence/test_postgres.rs @@ -0,0 +1,77 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Disposable `PostgreSQL` schemas for `#[ignore]`d backend tests. + +use super::Store; + +/// Names the disposable `PostgreSQL` database for `#[ignore]`d tests. +/// `mise run test:rust:postgres` always exports it. +pub const TEST_POSTGRES_URL_ENV: &str = "OPENSHELL_TEST_POSTGRES_URL"; + +/// A uniquely named schema. Stores connected through it see only that schema. +pub struct TestSchema { + admin: sqlx::PgPool, + schema: String, + url: String, +} + +impl TestSchema { + /// Create a schema named `_` in the database named by + /// [`TEST_POSTGRES_URL_ENV`]. + /// + /// Panics when the variable is unset; callers are `#[ignore]`d. + pub async fn create(prefix: &str) -> Self { + let base = std::env::var(TEST_POSTGRES_URL_ENV).unwrap_or_else(|_| { + panic!( + "{TEST_POSTGRES_URL_ENV} must name a disposable PostgreSQL database; \ + run mise run test:rust:postgres" + ) + }); + assert!( + base.starts_with("postgres"), + "{TEST_POSTGRES_URL_ENV} must be a postgres:// URL" + ); + let schema = format!("{prefix}_{}", uuid::Uuid::new_v4().simple()); + let admin = sqlx::PgPool::connect(&base) + .await + .expect("connect to the disposable PostgreSQL database"); + sqlx::query(sqlx::AssertSqlSafe(format!("CREATE SCHEMA {schema}"))) + .execute(&admin) + .await + .expect("create the test schema"); + let mut url = url::Url::parse(&base).expect("parse the PostgreSQL URL"); + url.query_pairs_mut() + .append_pair("options", &format!("-csearch_path={schema}")); + Self { + admin, + schema, + url: url.into(), + } + } + + /// Connection URL scoped to this schema. + pub fn url(&self) -> &str { + &self.url + } + + /// A new store with its own pool, as a separate gateway replica would + /// have. Runs migrations. + pub async fn connect_store(&self) -> Store { + Store::connect(&self.url) + .await + .expect("connect a store to the test schema") + } + + /// Drops only this test's schema. + pub async fn drop_schema(self) { + sqlx::query(sqlx::AssertSqlSafe(format!( + "DROP SCHEMA {} CASCADE", + self.schema + ))) + .execute(&self.admin) + .await + .expect("drop the test schema"); + self.admin.close().await; + } +} diff --git a/crates/openshell-server/src/provider_refresh.rs b/crates/openshell-server/src/provider_refresh.rs index c1b5184d18..a99d2560bf 100644 --- a/crates/openshell-server/src/provider_refresh.rs +++ b/crates/openshell-server/src/provider_refresh.rs @@ -1158,6 +1158,11 @@ async fn apply_minted_credential( credential_key: &str, minted: &MintedCredential, ) -> Result<(), Status> { + // Validate the expiration before staging anything, so this conversion can + // never leave staged credential handles behind. + let credential_expiration_time = + openshell_core::time::optional_timestamp_from_legacy_millis(minted.expires_at_ms) + .map_err(|error| Status::internal(error.to_string()))?; let mut updated = provider.clone(); let staging_id = format!("{}-refresh-{}", provider.object_id(), uuid::Uuid::new_v4()); let staged_handles = if let Some(credentials) = credentials @@ -1205,9 +1210,6 @@ async fn apply_minted_credential( } None }; - let credential_expiration_time = - openshell_core::time::optional_timestamp_from_legacy_millis(minted.expires_at_ms) - .map_err(|error| Status::internal(error.to_string()))?; if let Some(expiration_time) = credential_expiration_time.as_ref() { updated .credential_expiration_times @@ -1223,14 +1225,28 @@ async fn apply_minted_credential( updated.credential_expiration_times.remove(key); } } - // Acquire the shared sandbox mutation boundary only around validation and + // Acquire the workspace mutation key only around validation and // persistence, after any remote minting or credential staging. This // prevents route status from committing against the old provider revision // after the rotation writes, without holding the guard across network I/O. - let _sandbox_sync_guard = if let Some(compute) = compute { - Some(compute.sandbox_sync_guard().await.map_err(|error| { - Status::internal(format!("acquire provider mutation lock: {error}")) - })?) + let _mutation_guard = if let Some(compute) = compute { + match compute + .mutation_guard(crate::compute::MutationScope::Workspace(workspace)) + .await + { + Ok(guard) => Some(guard), + Err(error) => { + if let Some(credentials) = credentials + && let Some(handles) = &staged_handles + { + cleanup_staged_refresh_handles(credentials, provider, handles).await; + } + return Err(crate::grpc::persistence_error_to_status( + error, + "acquire provider mutation lock", + )); + } + } } else { None }; @@ -4130,6 +4146,139 @@ mod tests { assert_eq!(credentials.stored_credential_count(), Some(0)); } + #[tokio::test] + async fn apply_minted_credential_rejects_invalid_expiry_without_staging() { + use super::apply_minted_credential; + + let store = test_store().await; + let credentials = test_credentials(); + let mut prov = provider("expiring-aws", "aws"); + let original_handles = credentials + .store_provider_credentials( + prov.object_name(), + prov.object_workspace(), + prov.object_id(), + &HashMap::from([( + "AWS_ACCESS_KEY_ID".to_string(), + "old-access-key".to_string(), + )]), + &HashMap::new(), + ) + .await + .unwrap(); + prov.credential_handles.clone_from(&original_handles); + let stored_credential_count = credentials.stored_credential_count(); + store.put_message(&prov).await.unwrap(); + + // i64::MAX milliseconds is past the latest protobuf timestamp, so the + // expiration conversion fails. + let minted = super::MintedCredential { + access_token: "AKIAIOSFODNN7EXAMPLE".to_string(), + expires_at_ms: i64::MAX, + refresh_token: None, + additional_credentials: HashMap::new(), + }; + let err = apply_minted_credential( + &store, + "default", + Some(&credentials), + None, + &prov, + "AWS_ACCESS_KEY_ID", + &minted, + ) + .await + .unwrap_err(); + assert_eq!(err.code(), tonic::Code::Internal); + assert_eq!( + credentials.stored_credential_count(), + stored_credential_count + ); + let stored = store + .get_message_by_name::("default", "expiring-aws") + .await + .unwrap() + .unwrap(); + assert_eq!(stored.credential_handles, original_handles); + } + + #[tokio::test] + async fn apply_minted_credential_returns_unavailable_when_guard_times_out() { + use super::apply_minted_credential; + + let state = crate::grpc::test_support::test_server_state().await; + state + .compute + .set_mutation_lock_timeout_for_tests(std::time::Duration::from_millis(50)); + let credentials = test_credentials(); + let mut prov = provider("guarded-aws", "aws"); + let original_handles = credentials + .store_provider_credentials( + prov.object_name(), + prov.object_workspace(), + prov.object_id(), + &HashMap::from([( + "AWS_ACCESS_KEY_ID".to_string(), + "old-access-key".to_string(), + )]), + &HashMap::new(), + ) + .await + .unwrap(); + prov.credential_handles.clone_from(&original_handles); + let stored_credential_count = credentials.stored_credential_count(); + state.store.put_message(&prov).await.unwrap(); + let provider_writer = state + .compute + .mutation_guard(crate::compute::MutationScope::Workspace("default")) + .await + .unwrap(); + + let minted = super::MintedCredential { + access_token: "AKIAIOSFODNN7EXAMPLE".to_string(), + expires_at_ms: 4_000_000_000_000, + refresh_token: None, + additional_credentials: HashMap::new(), + }; + let err = apply_minted_credential( + &state.store, + "default", + Some(&credentials), + Some(&state.compute), + &prov, + "AWS_ACCESS_KEY_ID", + &minted, + ) + .await + .unwrap_err(); + assert_eq!(err.code(), tonic::Code::Unavailable); + let details = openshell_core::rpc_error::decode_details(&err).expect("error details"); + assert_eq!( + details.error_info().expect("error info").reason, + "MUTATION_LOCK_TIMEOUT" + ); + let stored = state + .store + .get_message_by_name::("default", "guarded-aws") + .await + .unwrap() + .unwrap(); + assert_eq!(stored.credential_handles, original_handles); + let resolved = credentials + .resolve_provider_handles(&stored, current_time_ms()) + .await + .unwrap(); + assert_eq!( + resolved.values.get("AWS_ACCESS_KEY_ID"), + Some(&"old-access-key".to_string()) + ); + assert_eq!( + credentials.stored_credential_count(), + stored_credential_count + ); + drop(provider_writer); + } + // A wiremock responder that blocks the STS response until the test releases // it, so a delete-refresh can be interleaved deterministically while the // rotation is parked awaiting STS. diff --git a/crates/openshell-server/src/ssh_identity.rs b/crates/openshell-server/src/ssh_identity.rs index 74e43c9ffd..224beaf3ce 100644 --- a/crates/openshell-server/src/ssh_identity.rs +++ b/crates/openshell-server/src/ssh_identity.rs @@ -307,7 +307,7 @@ impl SshIdentityStore { /// Keep credential staging and publication alive if the calling RPC is /// cancelled. The independent identity lock fences parent deletion on - /// every replica without nesting the creation cross-object lock. + /// every replica without nesting the creation mutation guard. pub(crate) async fn prepare( &self, sandbox: &mut Sandbox, diff --git a/crates/openshell-server/src/supervisor_session.rs b/crates/openshell-server/src/supervisor_session.rs index 494e0730f0..41f92cfbb8 100644 --- a/crates/openshell-server/src/supervisor_session.rs +++ b/crates/openshell-server/src/supervisor_session.rs @@ -2084,12 +2084,24 @@ async fn establish_supervisor_session( ) .await { - state + let was_current = state .supervisor_sessions - .remove_if_current(&sandbox_id, &session_id); + .remove_if_current(&sandbox_id, &session_id) + .is_some(); if let Err(err) = owner_index.release_if_current(&owner_guard).await { warn!(sandbox_id, session_id, error = %err, "supervisor session: failed to release owner after endpoint status initialization failure"); } + // While this session was current, the superseded session's disconnect + // reset deferred to this one, so its evidence may still be stored. + // Invalidate it the way a session end does. + if was_current { + tokio::spawn( + crate::grpc::policy::retry_endpoint_status_after_supervisor_disconnect( + Arc::clone(&state), + sandbox_id.clone(), + ), + ); + } return Err(error); } if !state diff --git a/deploy/helm/openshell/README.md b/deploy/helm/openshell/README.md index 4b534bc0be..f7e54677b0 100644 --- a/deploy/helm/openshell/README.md +++ b/deploy/helm/openshell/README.md @@ -340,7 +340,7 @@ discovery endpoint or its TLS CA. | agentSandbox.preflight.enabled | bool | `true` | Check the live cluster for a supported Agent Sandbox API before rendering gateway resources. Disable only for offline rendering and linting. | | autoscaling.behavior | object | `{"scaleDown":{"policies":[{"periodSeconds":120,"type":"Pods","value":1}],"stabilizationWindowSeconds":300}}` | HPA scaling behavior. Scale-down disconnects the removed pod's supervisor sessions; they reconnect to the remaining replicas, so the default removes at most one replica every two minutes after a five-minute stabilization window. Helm merges maps: set autoscaling.behavior.scaleDown to null to drop the default. | | autoscaling.enabled | bool | `false` | Render a HorizontalPodAutoscaler and stop rendering spec.replicas. | -| autoscaling.maxReplicas | int | `4` | Maximum gateway replicas. Each replica opens its own PostgreSQL connection pool; size the database for rollouts at this count, as the High Availability guide describes. | +| autoscaling.maxReplicas | int | `4` | Maximum gateway replicas. Each replica opens up to 14 PostgreSQL connections; size the database for rollouts at this count, as the High Availability guide describes. With a Deployment, the default of 4 needs `max_connections` of at least 126, more than the PostgreSQL default of 100. | | autoscaling.metrics | list | `[]` | Additional autoscaling/v2 MetricSpec entries appended verbatim, such as Pods metrics served by prometheus-adapter. | | autoscaling.minReplicas | int | `2` | Minimum gateway replicas. Use 2 or more to survive a pod failure. | | autoscaling.targetCPUUtilizationPercentage | int | `80` | Target average CPU utilization, as a percentage of resources.requests.cpu. Set to null to disable. Requires resources.requests.cpu, or resources.limits.cpu, which Kubernetes copies into the request. | diff --git a/deploy/helm/openshell/values.yaml b/deploy/helm/openshell/values.yaml index d0612cd563..d19fe56ed7 100644 --- a/deploy/helm/openshell/values.yaml +++ b/deploy/helm/openshell/values.yaml @@ -231,9 +231,10 @@ autoscaling: enabled: false # -- Minimum gateway replicas. Use 2 or more to survive a pod failure. minReplicas: 2 - # -- Maximum gateway replicas. Each replica opens its own PostgreSQL - # connection pool; size the database for rollouts at this count, as the High - # Availability guide describes. + # -- Maximum gateway replicas. Each replica opens up to 14 PostgreSQL + # connections; size the database for rollouts at this count, as the High + # Availability guide describes. With a Deployment, the default of 4 needs + # `max_connections` of at least 126, more than the PostgreSQL default of 100. maxReplicas: 4 # -- Target average CPU utilization, as a percentage of # resources.requests.cpu. Set to null to disable. Requires diff --git a/docs/kubernetes/high-availability.mdx b/docs/kubernetes/high-availability.mdx index 0153df050f..5de0e85891 100644 --- a/docs/kubernetes/high-availability.mdx +++ b/docs/kubernetes/high-availability.mdx @@ -204,6 +204,7 @@ configuration, and access control. | Relay rejections | `openshell_server_relay_rejected_total` | Alert on any increase. | | Relay claims | `openshell_server_relay_claim_duration_seconds`, `openshell_server_relay_expired_total` | Supervisor connect-back time for exec, SSH, forwarding, and service traffic, including relays requested through peers. A high 99th percentile while rejections stay flat points at the supervisor or its node, not at relay capacity. Relays not claimed within 10 seconds are missing from the histogram and count as expired, so alert on any increase. | | Routed requests | `openshell_server_routed_request_attempts_total`, `openshell_server_peer_request_duration_seconds` | Local/peer relay setup mix and outbound peer rate, failures, and latency. Counts completed or cancelled attempts, including retries. Spikes of `grpc_code="unavailable"` during rollouts are expected. Do not use as an HPA target. | +| Lock contention | `openshell_server_mutation_lock_wait_seconds`, `openshell_server_mutation_lock_timeouts_total` | Alert on any timeout. | These queries assume that Prometheus labels each series with `namespace` and `pod`: @@ -233,6 +234,12 @@ sum by (operation) (rate(openshell_server_routed_request_attempts_total{route="p # 99th percentile peer request latency across the fleet. histogram_quantile(0.99, sum by (le, operation) (rate(openshell_server_peer_request_duration_seconds_bucket{outcome="success"}[5m]))) + +# 99th percentile mutation lock wait by scope. Watch scope="sandbox" during rollouts. +histogram_quantile(0.99, sum by (le, scope) (rate(openshell_server_mutation_lock_wait_seconds_bucket[5m]))) + +# Mutation lock timeouts by scope. Alert on any increase. +sum by (scope) (increase(openshell_server_mutation_lock_timeouts_total[10m])) ``` Relay capacity is used on the replica that owns the sandbox's supervisor @@ -242,6 +249,16 @@ stays counted for up to about 40 seconds, the 10-second claim timeout plus the 30-second cleanup interval. Routed request counts include retries during rollouts. +Each supervisor that reconnects takes one or two short sandbox-scoped locks on +the replica that receives it. When a gateway pod stops during a rollout, its +supervisors reconnect to the other replicas, so watch +`openshell_server_mutation_lock_wait_seconds` and +`openshell_server_mutation_lock_timeouts_total` with `scope="sandbox"` on the +receiving replicas. Rising waits or any timeout mean that those replicas +receive reconnects faster than they can absorb them, or that slow compute +driver, credential backend, middleware, or provider profile source calls during +sandbox and provider changes hold their locks longer. + ## Scale the Gateway Scale the gateway by changing a fixed replica count or by letting a @@ -374,7 +391,9 @@ kubectl get --raw "/apis/custom.metrics.k8s.io/v1beta1/namespaces/openshell/pods ### Size PostgreSQL Connections -Each gateway pod opens up to 10 PostgreSQL connections on demand. The chart +Each gateway pod opens up to 14 PostgreSQL connections on demand, 10 for data +access and 4 for mutation locks. A request cancelled during a lock wait can +leave its lock session open for up to 10 seconds, so keep headroom. The chart uses the Kubernetes default rolling update strategy for each workload kind, so the number of pods that run during a rollout depends on the kind. @@ -384,18 +403,18 @@ and each replaced pod can keep its connections throughout its termination grace period. A rollout can therefore run up to twice the replica count at once, and an eviction or an autoscaler scale-down during the rollout adds more terminating pods. Set PostgreSQL `max_connections` to at least -`(2 × replicas + surge) × 10`, which leaves one surge of margin for those +`(2 × replicas + surge) × 14`, which leaves one surge of margin for those pods, plus headroom for your other clients and administration. Use `autoscaling.maxReplicas` as the replica count when autoscaling is enabled. | Replicas | Surge | Minimum `max_connections` | |---|---|---| -| 2 | 1 | `(4 + 1) × 10 = 50` | -| 3 | 1 | `(6 + 1) × 10 = 70` | -| 4 | 1 | `(8 + 1) × 10 = 90` | -| 6 | 2 | `(12 + 2) × 10 = 140` | +| 2 | 1 | `(4 + 1) × 14 = 70` | +| 3 | 1 | `(6 + 1) × 14 = 98` | +| 4 | 1 | `(8 + 1) × 14 = 126` | +| 6 | 2 | `(12 + 2) × 14 = 196` | -With a Deployment, five or more replicas exceed what the PostgreSQL default of +With a Deployment, three or more replicas exceed what the PostgreSQL default of 100 leaves for the gateway, because PostgreSQL reserves 3 connections for superusers. Raise `max_connections` or use a larger managed instance. Let one rollout finish before you start another, because each overlapping rollout adds @@ -404,7 +423,7 @@ its own terminating pods. If you run several replicas as a StatefulSet with `workload.allowMultiReplicaStatefulSet`, a rollout adds no surge pods. It replaces one pod at a time and creates each replacement only after the old pod -exits. Set `max_connections` to at least `replicas × 10`, using +exits. Set `max_connections` to at least `replicas × 14`, using `autoscaling.maxReplicas` when autoscaling is enabled, plus headroom for your other clients and administration. @@ -412,6 +431,36 @@ If a connection pooler sits between the gateway and PostgreSQL, use session pooling. The gateway holds session-level advisory locks, which transaction pooling breaks. +Mutations lock only what they change. Sandbox operations lock their own +sandbox, provider and workspace-profile changes lock their workspace, and +gateway-global policy, settings, and platform-profile changes lock the whole +fleet. Sandbox start, stop, and delete lock only on the replica that runs them, +and only a gateway-global change on that replica blocks them directly. +Provider, workspace-profile, and other replicas' gateway-global changes delay +them only while another operation on the same sandbox waits behind the change. +Sandbox create, start, restart, and delete also take one fleet-wide lock while +they provision or remove the sandbox's SSH host identity, so that step runs for +one sandbox at a time across all replicas. For other +mutations, a lock wait longer than 10 seconds returns `UNAVAILABLE` with the +reason `MUTATION_LOCK_TIMEOUT`, and clients can retry the request. A request +that carried a `request_id` leaves its admission unresolved, like any other +error, so a retry with the same ID returns `REQUEST_OUTCOME_UNCERTAIN`. Observe +resource state and reconcile effects, then send a new request with a new +`request_id`. Refer to +[Durable Request Admission](/sdk/api-errors#durable-request-admission). + +### Upgrade from an Earlier Release + +Each gateway pod now opens up to 14 PostgreSQL connections instead of 10. The +upgrade rollout already runs new pods at that count, so raise `max_connections` +to the minimum in [Size PostgreSQL Connections](#size-postgresql-connections) +before you upgrade. + +While gateways from the earlier release and the new release run together +during the upgrade, the older replicas serialize every mutation across the +fleet. Mutations can be slower and can briefly fail under load until the +rollout finishes. + ## Next Steps - To expose the gateway through a highly available data path, refer to diff --git a/docs/observability/gateway-metrics.mdx b/docs/observability/gateway-metrics.mdx index aa849977d3..d3291a3da7 100644 --- a/docs/observability/gateway-metrics.mdx +++ b/docs/observability/gateway-metrics.mdx @@ -3,7 +3,7 @@ # SPDX-License-Identifier: Apache-2.0 title: "Gateway Metrics" sidebar-title: "Gateway Metrics" -description: "Scrape Prometheus metrics from the OpenShell gateway, including supervisor session, relay capacity, and peer routing signals for multi-replica deployments." +description: "Scrape Prometheus metrics from the OpenShell gateway, including supervisor session, relay capacity, peer routing, and mutation lock signals for multi-replica deployments." keywords: "Generative AI, Cybersecurity, Observability, Metrics, Prometheus, Gateway, Kubernetes, High Availability" --- @@ -175,6 +175,13 @@ that owns a sandbox's supervisor session: | `openshell_server_routed_request_attempts_total` | Counter | `operation`, `route`, `relay_kind`, `outcome`, `grpc_code` | Local relay setup and outbound peer attempts, counted once when they finish or are cancelled. Each retry counts separately. | | `openshell_server_peer_request_duration_seconds` | Histogram | `operation`, `outcome` | Latency of outbound peer requests only (`route="peer"`). For relays, until the owner's supervisor claimed the relay. | +Mutation locks: + +| Metric | Type | Labels | Description | +|---|---|---|---| +| `openshell_server_mutation_lock_wait_seconds` | Histogram | `scope` | Time to acquire a mutation lock, recorded for successful acquisitions. | +| `openshell_server_mutation_lock_timeouts_total` | Counter | `scope` | Acquisitions that waited longer than 10 seconds, including background work such as supervisor disconnect handling and provider credential refresh. API requests fail with `UNAVAILABLE` and reason `MUTATION_LOCK_TIMEOUT`, and [API Errors](/sdk/api-errors) describes how to retry them. A lock connection that PostgreSQL does not open, although at least a second of the wait remained, is not counted. | + The labels take these values: - `operation` is `relay`, `report_provider_readiness`, `report_endpoint_status`, @@ -210,6 +217,7 @@ The labels take these values: unchanged. The label is not named `code` because `openshell_server_grpc_requests_total` uses that name for the numeric status code. +- `scope` is `global`, `workspace`, or `sandbox`. Gauges and counters exist from startup. They start at `0`, except `openshell_server_relay_pending_capacity`, which starts at its limit. `openshell_server_routed_request_attempts_total` starts with seven diff --git a/docs/sdk/api-errors.mdx b/docs/sdk/api-errors.mdx index 03cd9cf206..23a34c3b96 100644 --- a/docs/sdk/api-errors.mdx +++ b/docs/sdk/api-errors.mdx @@ -31,6 +31,7 @@ Recognized gateway reasons include the following. | `INVALID_ARGUMENT` | `INVALID_ARGUMENT` | Correct the fields listed in `BadRequest`. | | `RESOURCE_VERSION_CONFLICT` | `ABORTED` | Read the resource again and construct a new conditional write. `metadata.recovery` is `REFRESH_STATE`; `current_resource_version` is included when known. | | `PROFILE_SOURCE_UNAVAILABLE` | `UNAVAILABLE` | Retry a profile snapshot read after at least the supplied delay. | +| `MUTATION_LOCK_TIMEOUT` | `UNAVAILABLE` | The request waited more than 10 seconds for a concurrent mutation to release a sandbox, workspace, or gateway-wide lock. Retry after at least the supplied delay. A request that carried a `request_id` leaves its admission unresolved, like any other error, so a retry with that ID returns `REQUEST_OUTCOME_UNCERTAIN`. Observe resource state and reconcile effects before you start a new request with a new ID. | | `REQUEST_ID_PAYLOAD_MISMATCH` | `FAILED_PRECONDITION` | Keep the original payload for that request ID. Inspect the original operation before submitting a different one. | | `REQUEST_OUTCOME_UNCERTAIN` | `FAILED_PRECONDITION` | An attempt is admitted but has no confirmed replayable success. Observe resource state and reconcile effects. Do not switch to a new ID to bypass the admission. | | `REQUEST_REPLAY_UNAVAILABLE` | `FAILED_PRECONDITION` | The original scope, resource, interceptor transformation, or fingerprint key is no longer replayable. Reconcile effects; the gateway does not execute the request again. Missing private-key material can also reject admission before work starts. | diff --git a/skills/debug-openshell-cluster/SKILL.md b/skills/debug-openshell-cluster/SKILL.md index bd30ceb8ef..96b11cee37 100644 --- a/skills/debug-openshell-cluster/SKILL.md +++ b/skills/debug-openshell-cluster/SKILL.md @@ -471,18 +471,118 @@ kubectl -n get deployment,service,pod -l app.kubernetes.io/name= logs deployment/ --tail=200 ``` -Multi-replica gateways serialize cross-object sandbox and provider mutations -with a PostgreSQL advisory lock. If those RPCs stall while ordinary reads and -health checks remain responsive, inspect long-running database sessions and -advisory-lock waiters. Do not print the database URI or Secret contents into -logs: +Gateways that use an external PostgreSQL database guard cross-object +mutations with PostgreSQL advisory locks at three levels: one global key, one +key per workspace, and one key per sandbox, each taken shared or exclusive. +Sandbox operations on different sandboxes do not wait on each other's locks; +provider and workspace-profile changes block guarded sandbox mutations in +their workspace; gateway-global policy and settings changes block every +guarded mutation. Guarded sandbox mutations are sandbox create, provider +attach and detach, sandbox-scoped policy and settings changes, and supervisor +configuration, policy status, and endpoint status reports. Lifecycle paths +(start, stop, delete, restart, supervisor connection and process-exit state, +driver-watch updates, and reconcile) take only a process-local shared global +and exclusive sandbox lock, with no time limit and no mutation advisory lock +(start, restart, and delete also take the SSH identity key below), so only +a global change on the same replica blocks them directly. Provider, +workspace-profile, and other replicas' global changes delay them only while +another operation on the same sandbox waits behind the change, holding that +sandbox's key. Provisioning-deadline expiry also takes +its workspace key locally, so provider and workspace-profile changes on the +same replica can delay it. Each +replica takes its advisory locks on a dedicated pool of 4 PostgreSQL +connections, one per guard, so at most 4 guarded operations per replica hold +or wait for PostgreSQL locks at once. A request cancelled during its lock wait +frees its slot while its backend keeps waiting, with any keys it took, until +the 10-second deadline, so `pg_locks` can briefly show more lock sessions from +one pod. A guard that waits on a PostgreSQL lock +keeps its connection while it waits, so contention on one key can fill the +pool and make unrelated guarded operations on that replica queue for a +connection, within the same 10-second limit. + +A lock wait longer than 10 seconds returns `UNAVAILABLE` with reason +`MUTATION_LOCK_TIMEOUT` and the message "... timed out waiting for a +concurrent mutation; retry the request", logs +`mutation lock acquisition timed out` with `scope`, `waited_ms`, and `detail` +fields, and increments `openshell_server_mutation_lock_timeouts_total`. The +`detail` field says where the wait stopped: + +- `waiting for a local mutation lock`: another operation on the same replica + holds a conflicting key. +- `waiting for a mutation lock connection`: all 4 lock-pool connections of that + replica were in use, by waiters or by holders slowed by a compute driver, + credential backend, middleware, or provider profile source call, or a local + wait left less than a second to open one. `pg_locks` shows no row for the + timed-out operation. Look for the pod's granted or waiting advisory-lock + rows by `client_addr`, and find the holder they wait on. +- `waiting for a PostgreSQL advisory lock`, or PostgreSQL's own + `canceling statement due to lock timeout`: a conflicting key is held on + another connection, usually by another replica or by an older gateway during + an upgrade. `pg_locks` shows the holder and the waiters. + +A lock connection that PostgreSQL does not open, although at least a second of +the 10-second limit remained, is not a lock timeout. The request fails with +`INTERNAL` and "could not open a mutation lock connection", and the gateway +logs `mutation lock acquisition failed` without incrementing the timeout +counter. PostgreSQL refused the connection, ran out of connection slots, was +starting up, or did not answer. Check the PostgreSQL logs for "too many +clients already", "remaining connection slots are reserved", or "the database +system is starting up", and compare +`SELECT count(*) FROM pg_stat_activity` with `SHOW max_connections`. + +A gateway log line `timed out returning PostgreSQL mutation lock connection; +discarded connection` means returning a lock session stalled for 5 seconds and +the gateway dropped it; that slot stays busy for up to those 5 seconds. The +warning does not identify the PostgreSQL backend. Healthy guards also hold +idle advisory-lock sessions while validation and writes use separate data +connections. Neither an idle duration nor a matching `client_addr` proves +that a holder is orphaned, even when it blocks a timed-out request. + +Before using `SELECT pg_terminate_backend()`, conclusively map that +backend to its owning gateway process and confirm that process has stopped +or can no longer write. If the owner is still running, stop it first and +verify its exit; terminating or failing readiness alone is insufficient. Killing +its lock session while it can still write removes exclusion from an active +mutation. Recheck the backend PID and `backend_start` before terminating the +confirmed orphan: pod IPs and PIDs can be reused, and a database proxy can +hide several gateways behind one `client_addr`. If ownership cannot be +established, investigate connectivity instead of choosing a backend by IP or +idle state. Setting PostgreSQL +`tcp_keepalives_idle`, `tcp_keepalives_interval`, and `tcp_keepalives_count` +(for example 60, 10, and 6) bounds how long such sessions survive. + +If mutations stall or time out while reads and health checks work, inspect +advisory-lock holders and waiters. Do not print the database URI or Secret +contents into logs: ```sql -SELECT pid, granted, waitstart -FROM pg_locks -WHERE locktype = 'advisory'; +SELECT l.pid, l.mode, l.granted, l.waitstart, l.classid, l.objid, + a.client_addr, a.backend_start, a.state, + now() - a.state_change AS in_state_for +FROM pg_locks l +JOIN pg_stat_activity a USING (pid) +WHERE l.locktype = 'advisory' +ORDER BY l.granted, l.waitstart; ``` +Gateways hold these locks at session level outside any transaction, so a +holder (`granted = t`) usually shows `state = idle`. It still holds the lock, +and `in_state_for` approximates how long. `client_addr` is the holding +gateway pod's IP, or the pooler's IP when a connection pooler sits in between. + +The global key appears as `classid = 1330660686` and `objid = 1397247052` +(key `0x4F50454E53484C4C`). Older gateways during a rolling upgrade take that +key exclusively for every guarded mutation, so mutations can queue behind +them until the rollout finishes. + +SSH host identity provisioning and cleanup, in sandbox create, start, +restart, and delete, take one more fleet-wide key, `classid = 1330860872` and +`objid = 1213158228` (key `0x4F535348484F5354`), on a data connection instead +of the lock pool, with a 10-second PostgreSQL `lock_timeout`. Sandbox create +holds its mutation guard while it waits for that key. A wait that runs out +returns `UNAVAILABLE` with "lock sandbox SSH identity failed", and +`openshell_server_mutation_lock_timeouts_total` does not count it. + For multi-replica gateway installs, supervisor and client session traffic may be served by a non-owner gateway replica and relayed to the current supervisor owner over the internal `PeerRelay` RPC. Check the headless peer Service, @@ -528,7 +628,7 @@ for pod in $(kubectl -n openshell get pod \ -o jsonpath='{range .items[?(@.spec.containers[0].name=="openshell-gateway")]}{.metadata.name}{" "}{end}'); do echo "${pod}" kubectl get --raw "/api/v1/namespaces/openshell/pods/${pod}:9090/proxy/metrics" \ - | grep -E '^openshell_server_(supervisor_sessions|relay_pending|relay_rejected_total|routed_request_attempts_total)' + | grep -E '^openshell_server_(supervisor_sessions|relay_pending|relay_rejected_total|routed_request_attempts_total|mutation_lock_timeouts_total)' done kubectl -n openshell get hpa kubectl -n openshell describe hpa openshell @@ -550,7 +650,7 @@ kubectl -n openshell port-forward pod/ 9090:9090 >/dev/null & pf_pid=$! sleep 2 curl -s http://localhost:9090/metrics \ - | grep -E '^openshell_server_(supervisor_sessions|relay_pending|relay_rejected_total|routed_request_attempts_total)' + | grep -E '^openshell_server_(supervisor_sessions|relay_pending|relay_rejected_total|routed_request_attempts_total|mutation_lock_timeouts_total)' kill "${pf_pid}" ``` @@ -1005,6 +1105,8 @@ credential failures. | Kubernetes gateway pod pending | PVC unbound, taint, selector, or insufficient resources | `kubectl -n openshell describe pod ` | | Kubernetes sandbox pod stuck pending, workspace PVC unbound | Cluster has no default `StorageClass` and OpenShell does not set `storageClassName` on the workspace PVC (clusters with a default `StorageClass` bind fine without it) | `kubectl -n openshell describe pvc`; set `server.workspaceStorageClass` (gateway config `workspace_storage_class`) to a valid `StorageClass` | | Kubernetes gateway pod crash loops | Missing secret, bad DB URL, bad TLS config | `kubectl -n openshell logs deployment/openshell -c openshell-gateway` or `kubectl -n openshell logs statefulset/openshell -c openshell-gateway` | +| Mutating RPCs return `UNAVAILABLE` with "timed out waiting for a concurrent mutation" | Advisory-lock contention, a full 4-connection lock pool on one replica, a slow PostgreSQL, or older gateways still running during an upgrade | `openshell_server_mutation_lock_timeouts_total`, the `detail` field of `mutation lock acquisition timed out` in gateway logs, the `pg_locks` query in Step 6 | +| Mutating RPCs return `INTERNAL` with "could not open a mutation lock connection" | PostgreSQL out of connection slots, restarting, or unreachable | PostgreSQL logs, `SELECT count(*) FROM pg_stat_activity` against `max_connections`, the connection sizing in the High Availability guide | | `helm upgrade` fails with an `autoscaling.*` message | HPA values invalid: missing `resources.requests` (or `resources.limits`), `maxReplicas` above 1 without `server.externalDbSecret` (or on a StatefulSet without `workload.allowMultiReplicaStatefulSet`), no metric target, or min/max out of order. "`minReplicas` and `maxReplicas` are not set" means `--reuse-values` kept a release without the chart's autoscaling defaults | Fix the values named in the error; upgrade with `--reset-then-reuse-values` instead of `--reuse-values` | | HPA shows `` targets | No metrics-server for CPU/memory, or the metrics adapter does not serve the custom metric | `kubectl -n openshell describe hpa openshell`, `kubectl get --raw /apis/custom.metrics.k8s.io/v1beta1` | | One replica holds most sessions after a rollout | Expected: sessions stay where they reconnected | `openshell_server_supervisor_sessions` per pod; it fades as sandboxes are recreated | diff --git a/tasks/scripts/run-postgres-tests.sh b/tasks/scripts/run-postgres-tests.sh new file mode 100755 index 0000000000..ac303d7b4d --- /dev/null +++ b/tasks/scripts/run-postgres-tests.sh @@ -0,0 +1,83 @@ +#!/usr/bin/env bash +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +# Run the PostgreSQL-backed openshell-server tests: #[ignore] tests whose names +# start with `postgres_`. Point OPENSHELL_TEST_POSTGRES_URL at a disposable +# database, or leave it unset to start a throwaway PostgreSQL container with +# the local container engine (Docker or Podman). Never point it at a database +# that a running gateway uses: the tests take fleet-wide advisory locks. +# +# Extra arguments are passed to cargo nextest. Use test-name filters to run a +# subset, for example `postgres_concurrency`; a second -E filterset +# would widen the selection instead of narrowing it. + +set -euo pipefail + +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" +cd "${ROOT}" + +# Same pinned image as the Kubernetes e2e fixture (e2e/kubernetes/postgres-fixture.yaml). +POSTGRES_IMAGE="${OPENSHELL_TEST_POSTGRES_IMAGE:-mirror.gcr.io/library/postgres:17.10-alpine3.23@sha256:979c4379dd698aba0b890599a6104e082035f98ef31d9b9291ec22f2b13059ca}" +CONTAINER_NAME="" + +cleanup() { + if [ -n "${CONTAINER_NAME}" ]; then + ce rm -f "${CONTAINER_NAME}" >/dev/null 2>&1 || true + fi +} +trap cleanup EXIT + +start_postgres() { + local password port ready=0 + # shellcheck source=tasks/scripts/container-engine.sh + source "${ROOT}/tasks/scripts/container-engine.sh" + + password="$(od -An -N16 -tx1 /dev/urandom | tr -d ' \n')" + CONTAINER_NAME="openshell-test-postgres-$$" + echo "Starting disposable PostgreSQL (${POSTGRES_IMAGE})..." + # No --rm: the EXIT trap removes the container, and keeping it until then + # preserves its logs when PostgreSQL fails to start. + ce run -d --name "${CONTAINER_NAME}" \ + -e POSTGRES_USER=openshell \ + -e POSTGRES_PASSWORD="${password}" \ + -e POSTGRES_DB=openshell \ + -p 127.0.0.1::5432 \ + "${POSTGRES_IMAGE}" >/dev/null + + # The image's init phase runs a socket-only server, so a TCP probe succeeds + # only once the final server accepts connections. + for _ in $(seq 1 60); do + if ce exec "${CONTAINER_NAME}" pg_isready -h 127.0.0.1 -U openshell -d openshell >/dev/null 2>&1; then + ready=1 + break + fi + if [ "$(ce inspect -f '{{.State.Running}}' "${CONTAINER_NAME}" 2>/dev/null)" != "true" ]; then + break + fi + sleep 1 + done + if [ "${ready}" != "1" ]; then + echo "ERROR: PostgreSQL did not become ready within 60s or its container exited" >&2 + ce logs "${CONTAINER_NAME}" >&2 || true + exit 1 + fi + + port="$(ce port "${CONTAINER_NAME}" 5432/tcp | head -n1 | awk -F: '{print $NF}')" + echo "PostgreSQL is ready on 127.0.0.1:${port} (container ${CONTAINER_NAME})" + export OPENSHELL_TEST_POSTGRES_URL="postgres://openshell:${password}@127.0.0.1:${port}/openshell" +} + +if [ -z "${OPENSHELL_TEST_POSTGRES_URL:-}" ]; then + start_postgres +fi + +# The mutation-replay test predates OPENSHELL_TEST_POSTGRES_URL. Always use +# the selected database, even if a legacy URL is inherited from the caller. +export OPENSHELL_REPLAY_TEST_DATABASE_URL="${OPENSHELL_TEST_POSTGRES_URL}" +export OPENSHELL_TELEMETRY_ENABLED=false + +# Advisory locks are database-wide, so run the tests one at a time. +cargo nextest run -p openshell-server --features test-support \ + --run-ignored only --test-threads 1 \ + -E 'test(/(^|::)postgres_/)' "$@" diff --git a/tasks/scripts/test-postgres-test-runner.sh b/tasks/scripts/test-postgres-test-runner.sh new file mode 100755 index 0000000000..35d8f8b2ce --- /dev/null +++ b/tasks/scripts/test-postgres-test-runner.sh @@ -0,0 +1,33 @@ +#!/usr/bin/env bash +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +set -euo pipefail + +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" +TEST_TMP="$(mktemp -d)" +trap 'rm -rf "${TEST_TMP}"' EXIT + +# Capture the database passed to the legacy test without running Cargo or +# connecting to a database. An explicit URL also bypasses container startup. +cat >"${TEST_TMP}/cargo" <<'EOF' +#!/usr/bin/env bash +set -euo pipefail +printf '%s\n' "${OPENSHELL_REPLAY_TEST_DATABASE_URL:-}" >"${POSTGRES_RUNNER_TEST_RESULT}" +EOF +chmod +x "${TEST_TMP}/cargo" + +for legacy_url in "" "postgres://stale.example/other"; do + PATH="${TEST_TMP}:${PATH}" \ + OPENSHELL_TEST_POSTGRES_URL="postgres://selected.example/disposable" \ + OPENSHELL_REPLAY_TEST_DATABASE_URL="${legacy_url}" \ + POSTGRES_RUNNER_TEST_RESULT="${TEST_TMP}/database-url" \ + bash "${ROOT}/tasks/scripts/run-postgres-tests.sh" + + if [ "$(cat "${TEST_TMP}/database-url")" != "postgres://selected.example/disposable" ]; then + echo "FAIL: the PostgreSQL test runner did not use the selected database" >&2 + exit 1 + fi +done + +echo "PostgreSQL test runner database selection tests passed." diff --git a/tasks/test.toml b/tasks/test.toml index 8a90f3c084..bbbb279495 100644 --- a/tasks/test.toml +++ b/tasks/test.toml @@ -12,6 +12,7 @@ depends = [ "test:sbom", "test:install-sh", "test:build-env", + "test:postgres-runner", "test:gateway-pull-policy", "test:e2e-image-overrides", "test:gateway-config", @@ -108,6 +109,17 @@ run = [ run_windows = "powershell -NoProfile -ExecutionPolicy Bypass -File tasks/scripts/windows-msvc.ps1 test-precommit native" hide = true +["test:rust:postgres"] +description = "Run PostgreSQL-backed openshell-server tests against OPENSHELL_TEST_POSTGRES_URL or a disposable Docker/Podman PostgreSQL container" +run = "tasks/scripts/run-postgres-tests.sh" +run_windows = "echo Skipping test:rust:postgres: PostgreSQL integration tests need a Linux or macOS container engine." + +["test:postgres-runner"] +description = "Test PostgreSQL test runner database selection without a database or container engine" +run = "tasks/scripts/test-postgres-test-runner.sh" +run_windows = "echo Skipping test:postgres-runner: the Unix PostgreSQL test runner does not apply on Windows." +hide = true + ["test:python"] description = "Run Python tests" depends = ["python:proto"]