use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use parking_lot::Mutex;
use tokio::sync::Semaphore;
use tracing::{debug, info};
use crate::config::DuckDbConfig;
use crate::connection::DuckDbConnection;
use crate::error::{DuckDbError, DuckDbResult};
#[derive(Debug, Clone)]
pub struct PoolConfig {
pub max_connections: usize,
pub min_connections: usize,
pub connection_timeout_ms: u64,
}
impl Default for PoolConfig {
fn default() -> Self {
Self {
max_connections: 10,
min_connections: 1,
connection_timeout_ms: 30_000,
}
}
}
#[derive(Clone)]
pub struct DuckDbPool {
config: Arc<DuckDbConfig>,
pool_config: Arc<PoolConfig>,
connections: Arc<Mutex<Vec<DuckDbConnection>>>,
semaphore: Arc<Semaphore>,
}
impl DuckDbPool {
pub async fn new(config: DuckDbConfig) -> DuckDbResult<Self> {
Self::with_pool_config(config, PoolConfig::default()).await
}
pub async fn with_pool_config(
config: DuckDbConfig,
pool_config: PoolConfig,
) -> DuckDbResult<Self> {
let pool_config = if config.is_in_memory() {
PoolConfig {
max_connections: 1,
min_connections: 1,
..pool_config
}
} else {
pool_config
};
info!(
max_connections = pool_config.max_connections,
min_connections = pool_config.min_connections,
"Creating DuckDB connection pool"
);
let pool = Self {
config: Arc::new(config),
pool_config: Arc::new(pool_config.clone()),
connections: Arc::new(Mutex::new(Vec::new())),
semaphore: Arc::new(Semaphore::new(pool_config.max_connections)),
};
for _ in 0..pool_config.min_connections {
let conn = pool.create_connection()?;
pool.connections.lock().push(conn);
}
Ok(pool)
}
pub fn builder() -> DuckDbPoolBuilder {
DuckDbPoolBuilder::default()
}
pub async fn get(&self) -> DuckDbResult<PooledConnection> {
debug!("Acquiring connection from pool");
let timeout = Duration::from_millis(self.pool_config.connection_timeout_ms);
let permit = tokio::time::timeout(timeout, self.semaphore.clone().acquire_owned())
.await
.map_err(|_| {
DuckDbError::timeout(format!(
"timed out after {:?} waiting for a connection from the pool",
timeout
))
})?
.map_err(|e| DuckDbError::pool(format!("Failed to acquire semaphore: {}", e)))?;
let conn = {
let mut connections = self.connections.lock();
connections.pop()
};
let conn = match conn {
Some(c) => c,
None => self.create_connection()?,
};
Ok(PooledConnection {
conn: Some(conn),
pool: self.clone(),
poisoned: AtomicBool::new(false),
_permit: permit,
})
}
fn create_connection(&self) -> DuckDbResult<DuckDbConnection> {
debug!("Creating new DuckDB connection");
DuckDbConnection::new(&self.config)
}
fn return_connection(&self, conn: DuckDbConnection) {
let mut connections = self.connections.lock();
if connections.len() < self.pool_config.max_connections {
connections.push(conn);
}
}
pub fn status(&self) -> PoolStatus {
let available = self.connections.lock().len();
let permits = self.semaphore.available_permits();
PoolStatus {
max_connections: self.pool_config.max_connections,
available_connections: available,
available_permits: permits,
in_use: self.pool_config.max_connections - permits,
}
}
pub fn config(&self) -> &DuckDbConfig {
&self.config
}
pub fn pool_config(&self) -> &PoolConfig {
&self.pool_config
}
}
impl std::fmt::Debug for DuckDbPool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DuckDbPool")
.field("status", &self.status())
.finish()
}
}
#[derive(Debug, Clone)]
pub struct PoolStatus {
pub max_connections: usize,
pub available_connections: usize,
pub available_permits: usize,
pub in_use: usize,
}
pub struct PooledConnection {
conn: Option<DuckDbConnection>,
pool: DuckDbPool,
poisoned: AtomicBool,
_permit: tokio::sync::OwnedSemaphorePermit,
}
impl PooledConnection {
pub fn connection(&self) -> &DuckDbConnection {
self.conn.as_ref().expect("Connection already taken")
}
pub(crate) fn poison(&self) {
self.poisoned.store(true, Ordering::Release);
}
pub async fn query(
&self,
sql: &str,
params: &[prax_query::filter::FilterValue],
) -> DuckDbResult<Vec<serde_json::Value>> {
let conn = self.connection().clone();
let sql = sql.to_string();
let params = params.to_vec();
tokio::task::spawn_blocking(move || conn.query(&sql, ¶ms))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn query_one(
&self,
sql: &str,
params: &[prax_query::filter::FilterValue],
) -> DuckDbResult<serde_json::Value> {
let conn = self.connection().clone();
let sql = sql.to_string();
let params = params.to_vec();
tokio::task::spawn_blocking(move || conn.query_one(&sql, ¶ms))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn query_optional(
&self,
sql: &str,
params: &[prax_query::filter::FilterValue],
) -> DuckDbResult<Option<serde_json::Value>> {
let conn = self.connection().clone();
let sql = sql.to_string();
let params = params.to_vec();
tokio::task::spawn_blocking(move || conn.query_optional(&sql, ¶ms))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn query_rows(
&self,
sql: &str,
params: &[prax_query::filter::FilterValue],
) -> DuckDbResult<Vec<crate::row_ref::DuckDbRowRef>> {
let conn = self.connection().clone();
let sql = sql.to_string();
let params = params.to_vec();
tokio::task::spawn_blocking(move || conn.query_rows(&sql, ¶ms))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn execute(
&self,
sql: &str,
params: &[prax_query::filter::FilterValue],
) -> DuckDbResult<usize> {
let conn = self.connection().clone();
let sql = sql.to_string();
let params = params.to_vec();
tokio::task::spawn_blocking(move || conn.execute(&sql, ¶ms))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn execute_batch(&self, sql: &str) -> DuckDbResult<()> {
let conn = self.connection().clone();
let sql = sql.to_string();
tokio::task::spawn_blocking(move || conn.execute_batch(&sql))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn copy_to_parquet(&self, query: &str, path: &str) -> DuckDbResult<()> {
let conn = self.connection().clone();
let query = query.to_string();
let path = path.to_string();
tokio::task::spawn_blocking(move || conn.copy_to_parquet(&query, &path))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn copy_to_csv(&self, query: &str, path: &str, header: bool) -> DuckDbResult<()> {
let conn = self.connection().clone();
let query = query.to_string();
let path = path.to_string();
tokio::task::spawn_blocking(move || conn.copy_to_csv(&query, &path, header))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn query_parquet(&self, path: &str) -> DuckDbResult<Vec<serde_json::Value>> {
let conn = self.connection().clone();
let path = path.to_string();
tokio::task::spawn_blocking(move || conn.query_parquet(&path))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn query_csv(
&self,
path: &str,
header: bool,
) -> DuckDbResult<Vec<serde_json::Value>> {
let conn = self.connection().clone();
let path = path.to_string();
tokio::task::spawn_blocking(move || conn.query_csv(&path, header))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
pub async fn query_json(&self, path: &str) -> DuckDbResult<Vec<serde_json::Value>> {
let conn = self.connection().clone();
let path = path.to_string();
tokio::task::spawn_blocking(move || conn.query_json(&path))
.await
.map_err(|e| DuckDbError::internal(format!("Task join error: {}", e)))?
}
}
impl Drop for PooledConnection {
fn drop(&mut self) {
if let Some(conn) = self.conn.take() {
if !self.poisoned.load(Ordering::Acquire) {
self.pool.return_connection(conn);
}
}
}
}
impl std::fmt::Debug for PooledConnection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PooledConnection").finish_non_exhaustive()
}
}
#[derive(Debug, Default)]
pub struct DuckDbPoolBuilder {
config: Option<DuckDbResult<DuckDbConfig>>,
pool_config: PoolConfig,
}
impl DuckDbPoolBuilder {
pub fn new() -> Self {
Self::default()
}
pub fn config(mut self, config: DuckDbConfig) -> Self {
self.config = Some(Ok(config));
self
}
pub fn path(mut self, path: &str) -> Self {
self.config = Some(DuckDbConfig::from_path(path));
self
}
pub fn in_memory(mut self) -> Self {
self.config = Some(Ok(DuckDbConfig::in_memory()));
self
}
pub fn url(mut self, url: &str) -> Self {
self.config = Some(DuckDbConfig::from_url(url));
self
}
pub fn max_connections(mut self, max: usize) -> Self {
self.pool_config.max_connections = max;
self
}
pub fn min_connections(mut self, min: usize) -> Self {
self.pool_config.min_connections = min;
self
}
pub fn connection_timeout_ms(mut self, timeout: u64) -> Self {
self.pool_config.connection_timeout_ms = timeout;
self
}
pub async fn build(self) -> DuckDbResult<DuckDbPool> {
let config = self
.config
.ok_or_else(|| DuckDbError::config("Database configuration required"))??;
DuckDbPool::with_pool_config(config, self.pool_config).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_pool_creation() {
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let status = pool.status();
assert_eq!(status.max_connections, 1);
assert!(status.available_connections >= 1);
}
#[tokio::test]
async fn test_pool_get_connection() {
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let conn = pool.get().await.unwrap();
let results = conn.query("SELECT 1 as value", &[]).await.unwrap();
assert_eq!(results.len(), 1);
}
#[tokio::test]
async fn test_pool_builder() {
let pool = DuckDbPool::builder()
.in_memory()
.max_connections(5)
.min_connections(2)
.build()
.await
.unwrap();
let status = pool.status();
assert_eq!(status.max_connections, 1);
assert_eq!(status.available_connections, 1);
}
#[tokio::test]
async fn test_in_memory_pool_shares_single_database() {
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
{
let conn = pool.get().await.unwrap();
conn.execute("CREATE TABLE shared_writes (value INTEGER)", &[])
.await
.unwrap();
conn.execute("INSERT INTO shared_writes VALUES (42)", &[])
.await
.unwrap();
}
{
let conn = pool.get().await.unwrap();
let rows = conn
.query("SELECT value FROM shared_writes", &[])
.await
.unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0]["value"], serde_json::json!(42));
}
}
#[tokio::test]
async fn test_connection_returned_to_pool() {
let pool = DuckDbPool::builder()
.in_memory()
.max_connections(2)
.min_connections(0)
.build()
.await
.unwrap();
let initial_permits = pool.semaphore.available_permits();
{
let _conn = pool.get().await.unwrap();
assert_eq!(pool.semaphore.available_permits(), initial_permits - 1);
}
assert_eq!(pool.semaphore.available_permits(), initial_permits);
}
#[tokio::test]
async fn test_pool_get_times_out_when_saturated() {
let pool = DuckDbPool::builder()
.in_memory()
.connection_timeout_ms(50)
.build()
.await
.unwrap();
let held = pool.get().await.unwrap();
let err = pool.get().await.unwrap_err();
assert!(
matches!(err, DuckDbError::Timeout(_)),
"expected a timeout error, got: {err:?}"
);
assert!(err.to_string().contains("timed out after 50ms"));
drop(held);
pool.get().await.unwrap();
}
}