Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
134 changes: 122 additions & 12 deletions crates/pgls_workspace/src/workspace/server/connection_manager.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
use std::collections::HashMap;
use std::hash::{DefaultHasher, Hash, Hasher};
use std::str::FromStr;
use std::sync::RwLock;
use std::time::{Duration, Instant};
Expand All @@ -12,15 +13,28 @@ use super::connection_key::ConnectionKey;
const INITIAL_FAILURE_BACKOFF: Duration = Duration::from_secs(5);
const MAX_FAILURE_BACKOFF: Duration = Duration::from_secs(60);

/// Fingerprints the connection parameters that [`ConnectionKey`] does not cover, so that a change
/// to one of them invalidates the cached pool instead of being silently ignored. Hashing keeps the
/// credentials out of the cache entries, and therefore out of anything that prints them.
fn settings_fingerprint(settings: &DatabaseSettings) -> u64 {
let mut hasher = DefaultHasher::new();
settings.connection_string.hash(&mut hasher);
settings.password.hash(&mut hasher);
settings.conn_timeout_secs.hash(&mut hasher);
hasher.finish()
}

/// Cached connection pool with last access time
struct CachedPool {
pool: PgPool,
settings_fingerprint: u64,
last_accessed: Instant,
idle_timeout: Duration,
}

struct CachedFailure {
message: String,
settings_fingerprint: u64,
attempts: u32,
next_retry_at: Instant,
}
Expand Down Expand Up @@ -51,20 +65,30 @@ impl ConnectionManager {
return None;
}

if self.connection_is_in_backoff(&key) {
let settings_fingerprint = settings_fingerprint(settings);
if self.connection_is_in_backoff(&key, settings_fingerprint) {
return None;
}

{
if let Ok(pools) = self.pools.read()
&& let Some(cached_pool) = pools.get(&key)
&& cached_pool.settings_fingerprint == settings_fingerprint
{
return Some(cached_pool.pool.clone());
}
}

let mut pools = self.pools.write().unwrap();

// Drop a pool created with stale connection parameters.
if pools
.get(&key)
.is_some_and(|pool| pool.settings_fingerprint != settings_fingerprint)
{
pools.remove(&key);
}

// Double-check after acquiring write lock
if let Some(cached_pool) = pools.get_mut(&key) {
cached_pool.last_accessed = Instant::now();
Expand Down Expand Up @@ -113,6 +137,7 @@ impl ConnectionManager {

let cached_pool = CachedPool {
pool: pool.clone(),
settings_fingerprint,
last_accessed: Instant::now(),
// TODO: add this to the db settings, for now default to five minutes
idle_timeout: Duration::from_secs(60 * 5),
Expand All @@ -129,38 +154,44 @@ impl ConnectionManager {
operation: impl FnOnce(&PgPool) -> Result<T, WorkspaceError>,
) -> Option<Result<T, WorkspaceError>> {
let pool = self.get_pool(settings)?;
let key = ConnectionKey::from(settings);
let result = operation(&pool);
self.record_result(&pool, &result);
self.record_result(&key, settings_fingerprint(settings), &result);
Some(result)
}

fn record_result<T>(&self, pool: &PgPool, result: &Result<T, WorkspaceError>) {
fn record_result<T>(
&self,
key: &ConnectionKey,
settings_fingerprint: u64,
result: &Result<T, WorkspaceError>,
) {
match result {
Ok(_) => self.clear_failure(&pool.into()),
Ok(_) => self.clear_failure(key),
Err(err @ WorkspaceError::DatabaseConnectionError(_)) => {
self.record_failure(pool, &err.to_string());
self.record_failure(key, settings_fingerprint, &err.to_string());
}
Err(_) => {}
}
}

fn record_failure(&self, pool: &PgPool, error: &str) {
let key: ConnectionKey = pool.into();
fn record_failure(&self, key: &ConnectionKey, settings_fingerprint: u64, error: &str) {
let mut failures = self.failures.write().unwrap();
let now = Instant::now();
let attempts = failures.get(&key).map_or(1, |failure| failure.attempts + 1);
let attempts = failures.get(key).map_or(1, |failure| failure.attempts + 1);
let multiplier = 1u32
.checked_shl(attempts.saturating_sub(1))
.unwrap_or(u32::MAX);
let backoff = INITIAL_FAILURE_BACKOFF
.saturating_mul(multiplier)
.min(MAX_FAILURE_BACKOFF);

let was_cached = failures.contains_key(&key);
let was_cached = failures.contains_key(key);
failures.insert(
key,
key.clone(),
CachedFailure {
message: error.to_string(),
settings_fingerprint,
attempts,
next_retry_at: now + backoff,
},
Expand All @@ -179,8 +210,15 @@ impl ConnectionManager {
}
}

fn connection_is_in_backoff(&self, key: &ConnectionKey) -> bool {
let failures = self.failures.read().unwrap();
fn connection_is_in_backoff(&self, key: &ConnectionKey, settings_fingerprint: u64) -> bool {
let mut failures = self.failures.write().unwrap();
if failures
.get(key)
.is_some_and(|failure| failure.settings_fingerprint != settings_fingerprint)
{
failures.remove(key);
return false;
}
let Some(failure) = failures.get(key) else {
return false;
};
Expand All @@ -202,3 +240,75 @@ impl ConnectionManager {
self.failures.write().unwrap().remove(key);
}
}

#[cfg(test)]
mod tests {
use super::*;

fn settings() -> DatabaseSettings {
DatabaseSettings {
enable_connection: true,
..DatabaseSettings::default()
}
}

#[tokio::test]
async fn same_settings_reuse_cached_pool() {
let manager = ConnectionManager::new();
let settings = settings();
let first = manager.get_pool(&settings).unwrap();

first.close().await;
let second = manager.get_pool(&settings).unwrap();

assert!(second.is_closed());
assert_eq!(manager.pools.read().unwrap().len(), 1);
}

#[tokio::test]
async fn password_change_replaces_cached_pool_with_new_settings() {
let manager = ConnectionManager::new();
let mut settings = settings();
let first_fingerprint = settings_fingerprint(&settings);
let first = manager.get_pool(&settings).unwrap();
first.close().await;

settings.password = "changed-password".to_string();
let changed_fingerprint = settings_fingerprint(&settings);
let pool = manager.get_pool(&settings).unwrap();

assert!(!pool.is_closed());
let pools = manager.pools.read().unwrap();
let cached = pools.get(&ConnectionKey::from(&settings)).unwrap();
assert_ne!(first_fingerprint, changed_fingerprint);
assert_eq!(cached.settings_fingerprint, changed_fingerprint);
assert_eq!(
cached.pool.connect_options().get_host(),
pool.connect_options().get_host()
);
assert_eq!(pools.len(), 1);
}

#[test]
fn password_change_clears_failure_backoff() {
let manager = ConnectionManager::new();
let mut settings = settings();
let key = ConnectionKey::from(&settings);
let old_fingerprint = settings_fingerprint(&settings);
manager.failures.write().unwrap().insert(
key.clone(),
CachedFailure {
message: "authentication failed".to_string(),
settings_fingerprint: old_fingerprint,
attempts: 1,
next_retry_at: Instant::now() + Duration::from_secs(60),
},
);

settings.password = "corrected-password".to_string();
let new_fingerprint = settings_fingerprint(&settings);
assert!(!manager.connection_is_in_backoff(&key, new_fingerprint));
assert!(!manager.failures.read().unwrap().contains_key(&key));
assert!(manager.get_pool(&settings).is_some());
}
}
Loading