fix: avoid lock contention for native workers on cached connection (#5481)

This commit is contained in:
Ruben Fiszel
2025-03-14 11:58:03 +01:00
committed by GitHub
parent 09a2791e2e
commit 8e95bc3972
+92 -52
View File
@@ -1,6 +1,6 @@
use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
@@ -17,7 +17,7 @@ use serde::Deserialize;
use serde_json::value::RawValue;
use serde_json::Map;
use serde_json::Value;
use tokio::sync::Mutex;
use tokio::sync::{Mutex, RwLock};
use tokio_postgres::Client;
use tokio_postgres::{types::ToSql, NoTls, Row};
use tokio_postgres::{
@@ -55,8 +55,9 @@ struct PgDatabase {
lazy_static! {
pub static ref CONNECTION_CACHE: Arc<Mutex<Option<(String, tokio_postgres::Client)>>> =
Arc::new(Mutex::new(None));
pub static ref CONNECTION_COUNTER: Arc<RwLock<HashMap<String, u64>>> =
Arc::new(RwLock::new(HashMap::new()));
pub static ref LAST_QUERY: AtomicU64 = AtomicU64::new(0);
pub static ref RUNNING: AtomicBool = AtomicBool::new(false);
}
fn do_postgresql_inner<'a>(
@@ -211,14 +212,10 @@ pub async fn do_postgresql(
);
let database_string_clone = database_string.clone();
RUNNING.store(true, std::sync::atomic::Ordering::Relaxed);
LAST_QUERY.store(
chrono::Utc::now().timestamp().try_into().unwrap_or(0),
std::sync::atomic::Ordering::Relaxed,
);
let mtex;
if !*CLOUD_HOSTED {
mtex = Some(CONNECTION_CACHE.lock().await);
mtex = CONNECTION_CACHE.try_lock().ok();
increment_connection_counter(&database_string).await;
} else {
mtex = None;
}
@@ -226,9 +223,15 @@ pub async fn do_postgresql(
let has_cached_con = mtex
.as_ref()
.is_some_and(|x| x.as_ref().is_some_and(|y| y.0 == database_string));
let new_client = if has_cached_con {
// tracing::error!("HAS CACHED CON: {}", has_cached_con);
let (new_client, mtex) = if has_cached_con {
tracing::info!("Using cached connection");
None
LAST_QUERY.store(
chrono::Utc::now().timestamp().try_into().unwrap_or(0),
std::sync::atomic::Ordering::Relaxed,
);
(None, mtex)
} else if sslmode == "require" {
tracing::info!("Creating new connection");
let mut connector = TlsConnector::builder();
@@ -266,7 +269,7 @@ pub async fn do_postgresql(
tracing::error!("connection error: {}", e);
}
});
Some((client, handle))
(Some((client, handle)), None)
} else {
tracing::info!("Creating new connection");
let (client, connection) = tokio::time::timeout(
@@ -284,7 +287,7 @@ pub async fn do_postgresql(
tracing::error!("connection error: {}", e);
}
});
Some((client, handle))
(Some((client, handle)), None)
};
let queries = parse_sql_blocks(query);
@@ -359,53 +362,77 @@ pub async fn do_postgresql(
)
.await?;
// drop the mtex to avoid holding the lock for too long, result has been returned
drop(mtex);
*mem_peak = size.load(Ordering::Relaxed) as i32;
RUNNING.store(false, std::sync::atomic::Ordering::Relaxed);
if let Some(handle) = handle {
if let Some(mut mtex) = mtex {
let abort_handler = handle.abort_handle();
if !*CLOUD_HOSTED {
// tracing::error!("Found handle");
if let Ok(mut mtex) = CONNECTION_CACHE.try_lock() {
if mtex.as_ref().is_none_or(|x| x.0 != database_string) {
// tracing::error!("Locked conn cached");
let abort_handler = handle.abort_handle();
if let Some(new_client) = new_client {
*mtex = Some((database_string, new_client.0));
}
drop(mtex);
LAST_QUERY.store(
chrono::Utc::now().timestamp().try_into().unwrap_or(0),
std::sync::atomic::Ordering::Relaxed,
);
tokio::spawn(async move {
loop {
tokio::time::sleep(Duration::from_secs(5)).await;
let last_query = LAST_QUERY.load(std::sync::atomic::Ordering::Relaxed);
let now = chrono::Utc::now().timestamp().try_into().unwrap_or(0);
//we cache connection for 5 minutes at most
if last_query + 60 * 5 < now
&& !RUNNING.load(std::sync::atomic::Ordering::Relaxed)
{
tracing::info!("Closing cache connection due to inactivity");
break;
}
let mtex = CONNECTION_CACHE.lock().await;
if mtex.is_none() {
// connection is not in the mutex anymore
break;
} else if let Some(mtex) = mtex.as_ref() {
if mtex.0.as_str() != &database_string_clone {
// connection is not the latest one
break;
let mut cache_new_con = false;
if let Some(new_client) = new_client {
cache_new_con = is_most_used_conn(&database_string).await;
if cache_new_con {
*mtex = Some((database_string, new_client.0));
} else {
new_client.1.abort();
}
} else {
handle.abort();
}
tracing::debug!("Keeping cached connection alive due to activity")
LAST_QUERY.store(
chrono::Utc::now().timestamp().try_into().unwrap_or(0),
std::sync::atomic::Ordering::Relaxed,
);
if cache_new_con {
tokio::spawn(async move {
loop {
tokio::time::sleep(Duration::from_secs(5)).await;
let last_query =
LAST_QUERY.load(std::sync::atomic::Ordering::Relaxed);
let now = chrono::Utc::now().timestamp().try_into().unwrap_or(0);
//we cache connection for 5 minutes at most
if last_query + 60 * 1 < now {
// tracing::error!("Closing cache connection due to inactivity");
tracing::info!(
"Closing cache pg executor connection due to inactivity"
);
break;
}
let mtex = CONNECTION_CACHE.lock().await;
if mtex.is_none() {
// connection is not in the mutex anymore
break;
} else if let Some(mtex) = mtex.as_ref() {
if mtex.0.as_str() != &database_string_clone {
// connection is not the latest one
break;
}
}
tracing::debug!(
"Keeping cached pg executor connection alive due to activity"
)
}
let mut mtex = CONNECTION_CACHE.lock().await;
*mtex = None;
abort_handler.abort();
});
}
} else {
handle.abort();
}
let mut mtex = CONNECTION_CACHE.lock().await;
*mtex = None;
abort_handler.abort();
});
} else {
handle.abort();
}
} else {
handle.abort();
}
@@ -416,6 +443,18 @@ pub async fn do_postgresql(
return Ok(raw_result);
}
async fn is_most_used_conn(database_string: &str) -> bool {
let counter_map = CONNECTION_COUNTER.read().await;
let current_count = counter_map.get(database_string).copied().unwrap_or(0);
let max_count = counter_map.values().copied().max().unwrap_or(0);
current_count >= max_count
}
async fn increment_connection_counter(database_string: &str) {
let mut counter_map = CONNECTION_COUNTER.write().await;
*counter_map.entry(database_string.to_string()).or_insert(0) += 1;
}
fn map_as_single_type<T>(
vec: Option<&Vec<Value>>,
f: impl Fn(&Value) -> Option<T>,
@@ -766,6 +805,7 @@ pub fn pg_cell_to_json_value(
Type::BYTEA_ARRAY => get_array(row, column, column_i, |a: Vec<u8>| {
Ok(JSONValue::String(format!("\\x{}", hex::encode(a))))
})?,
Type::VOID => JSONValue::Null,
_ => get_basic(row, column, column_i, |a: String| Ok(JSONValue::String(a)))?,
})
}