use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use serde_json::Value as JsonValue;
use tracing::{debug, instrument};
use prax_query::filter::FilterValue;
use prax_query::types::SortOrder;
use crate::error::{DuckDbError, DuckDbResult};
use crate::pool::{DuckDbPool, PooledConnection};
#[derive(Clone)]
pub struct DuckDbEngine {
pool: DuckDbPool,
tx_conn: Option<Arc<PooledConnection>>,
tx_finalized: Option<Arc<AtomicBool>>,
}
enum ConnectionSource<'a> {
Tx(&'a PooledConnection),
Pool(PooledConnection),
}
impl std::ops::Deref for ConnectionSource<'_> {
type Target = PooledConnection;
fn deref(&self) -> &PooledConnection {
match self {
Self::Tx(conn) => conn,
Self::Pool(conn) => conn,
}
}
}
const TX_FINALIZED: &str = "transaction has already been committed or rolled back";
#[derive(Debug, Clone)]
pub struct DuckDbQueryResult {
pub data: JsonValue,
}
impl DuckDbQueryResult {
pub fn new(data: JsonValue) -> Self {
Self { data }
}
pub fn json(&self) -> &JsonValue {
&self.data
}
pub fn into_json(self) -> JsonValue {
self.data
}
}
impl DuckDbEngine {
pub fn new(pool: DuckDbPool) -> Self {
Self {
pool,
tx_conn: None,
tx_finalized: None,
}
}
pub fn pool(&self) -> &DuckDbPool {
&self.pool
}
async fn connection(&self) -> DuckDbResult<ConnectionSource<'_>> {
if let Some(tx) = &self.tx_conn {
if self
.tx_finalized
.as_ref()
.is_some_and(|f| f.load(Ordering::Acquire))
{
return Err(DuckDbError::internal(TX_FINALIZED));
}
Ok(ConnectionSource::Tx(tx.as_ref()))
} else {
self.pool.get().await.map(ConnectionSource::Pool)
}
}
fn build_select(
&self,
table: &str,
columns: &[String],
filters: &HashMap<String, FilterValue>,
sort: &[(String, SortOrder)],
limit: Option<u64>,
offset: Option<u64>,
) -> (String, Vec<FilterValue>) {
let mut sql = String::new();
let mut params: Vec<FilterValue> = Vec::new();
let cols = if columns.is_empty() {
"*".to_string()
} else {
columns
.iter()
.map(|c| format!("\"{}\"", c))
.collect::<Vec<_>>()
.join(", ")
};
sql.push_str(&format!("SELECT {} FROM \"{}\"", cols, table));
if !filters.is_empty() {
let mut conditions = Vec::new();
for (field, value) in filters {
match value {
FilterValue::Null => {
conditions.push(format!("\"{}\" IS NULL", field));
}
_ => {
conditions.push(format!("\"{}\" = ?", field));
params.push(value.clone());
}
}
}
sql.push_str(" WHERE ");
sql.push_str(&conditions.join(" AND "));
}
if !sort.is_empty() {
let order_parts: Vec<String> = sort
.iter()
.map(|(col, dir)| {
let direction = match dir {
SortOrder::Asc => "ASC",
SortOrder::Desc => "DESC",
};
format!("\"{}\" {}", col, direction)
})
.collect();
sql.push_str(" ORDER BY ");
sql.push_str(&order_parts.join(", "));
}
if let Some(lim) = limit {
sql.push_str(&format!(" LIMIT {}", lim));
}
if let Some(off) = offset {
sql.push_str(&format!(" OFFSET {}", off));
}
(sql, params)
}
fn build_insert(
&self,
table: &str,
data: &HashMap<String, FilterValue>,
) -> (String, Vec<FilterValue>) {
let mut columns = Vec::new();
let mut placeholders = Vec::new();
let mut params: Vec<FilterValue> = Vec::new();
for (col, val) in data {
columns.push(format!("\"{}\"", col));
placeholders.push("?".to_string());
params.push(val.clone());
}
let sql = format!(
"INSERT INTO \"{}\" ({}) VALUES ({})",
table,
columns.join(", "),
placeholders.join(", ")
);
(sql, params)
}
fn build_update(
&self,
table: &str,
data: &HashMap<String, FilterValue>,
filters: &HashMap<String, FilterValue>,
) -> (String, Vec<FilterValue>) {
let mut params: Vec<FilterValue> = Vec::new();
let set_parts: Vec<String> = data
.iter()
.map(|(col, val)| {
params.push(val.clone());
format!("\"{}\" = ?", col)
})
.collect();
let mut sql = format!("UPDATE \"{}\" SET {}", table, set_parts.join(", "));
if !filters.is_empty() {
let mut conditions = Vec::new();
for (field, value) in filters {
match value {
FilterValue::Null => {
conditions.push(format!("\"{}\" IS NULL", field));
}
_ => {
conditions.push(format!("\"{}\" = ?", field));
params.push(value.clone());
}
}
}
sql.push_str(" WHERE ");
sql.push_str(&conditions.join(" AND "));
}
(sql, params)
}
fn build_delete(
&self,
table: &str,
filters: &HashMap<String, FilterValue>,
) -> (String, Vec<FilterValue>) {
let mut sql = format!("DELETE FROM \"{}\"", table);
let mut params: Vec<FilterValue> = Vec::new();
if !filters.is_empty() {
let mut conditions = Vec::new();
for (field, value) in filters {
match value {
FilterValue::Null => {
conditions.push(format!("\"{}\" IS NULL", field));
}
_ => {
conditions.push(format!("\"{}\" = ?", field));
params.push(value.clone());
}
}
}
sql.push_str(" WHERE ");
sql.push_str(&conditions.join(" AND "));
}
(sql, params)
}
#[instrument(skip(self, columns, filters, sort), fields(table = %table))]
pub async fn query_many(
&self,
table: &str,
columns: &[String],
filters: &HashMap<String, FilterValue>,
sort: &[(String, SortOrder)],
limit: Option<u64>,
offset: Option<u64>,
) -> DuckDbResult<Vec<DuckDbQueryResult>> {
let (sql, params) = self.build_select(table, columns, filters, sort, limit, offset);
debug!(sql = %sql, "Executing query_many");
let conn = self.connection().await?;
let results = conn.query(&sql, ¶ms).await?;
Ok(results.into_iter().map(DuckDbQueryResult::new).collect())
}
#[instrument(skip(self, columns, filters), fields(table = %table))]
pub async fn query_one(
&self,
table: &str,
columns: &[String],
filters: &HashMap<String, FilterValue>,
) -> DuckDbResult<DuckDbQueryResult> {
let (sql, params) = self.build_select(table, columns, filters, &[], Some(1), None);
debug!(sql = %sql, "Executing query_one");
let conn = self.connection().await?;
let result = conn.query_one(&sql, ¶ms).await?;
Ok(DuckDbQueryResult::new(result))
}
#[instrument(skip(self, columns, filters), fields(table = %table))]
pub async fn query_optional(
&self,
table: &str,
columns: &[String],
filters: &HashMap<String, FilterValue>,
) -> DuckDbResult<Option<DuckDbQueryResult>> {
let (sql, params) = self.build_select(table, columns, filters, &[], Some(1), None);
debug!(sql = %sql, "Executing query_optional");
let conn = self.connection().await?;
let result = conn.query_optional(&sql, ¶ms).await?;
Ok(result.map(DuckDbQueryResult::new))
}
#[instrument(skip(self, data), fields(table = %table))]
pub async fn execute_insert(
&self,
table: &str,
data: &HashMap<String, FilterValue>,
) -> DuckDbResult<DuckDbQueryResult> {
let (sql, params) = self.build_insert(table, data);
let sql = format!("{} RETURNING *", sql);
debug!(sql = %sql, "Executing insert");
let conn = self.connection().await?;
let rows = conn.query(&sql, ¶ms).await?;
let row = rows
.into_iter()
.next()
.ok_or_else(|| DuckDbError::query("INSERT ... RETURNING produced no row"))?;
Ok(DuckDbQueryResult::new(row))
}
#[instrument(skip(self, data, filters), fields(table = %table))]
pub async fn execute_update(
&self,
table: &str,
data: &HashMap<String, FilterValue>,
filters: &HashMap<String, FilterValue>,
) -> DuckDbResult<u64> {
let (sql, params) = self.build_update(table, data, filters);
debug!(sql = %sql, "Executing update");
let conn = self.connection().await?;
let affected = conn.execute(&sql, ¶ms).await?;
Ok(affected as u64)
}
#[instrument(skip(self, filters), fields(table = %table))]
pub async fn execute_delete(
&self,
table: &str,
filters: &HashMap<String, FilterValue>,
) -> DuckDbResult<u64> {
let (sql, params) = self.build_delete(table, filters);
debug!(sql = %sql, "Executing delete");
let conn = self.connection().await?;
let affected = conn.execute(&sql, ¶ms).await?;
Ok(affected as u64)
}
#[instrument(skip(self, params), fields(sql = %sql))]
pub async fn execute_raw(
&self,
sql: &str,
params: &[FilterValue],
) -> DuckDbResult<Vec<DuckDbQueryResult>> {
debug!("Executing raw SQL");
let conn = self.connection().await?;
let results = conn.query(sql, params).await?;
Ok(results.into_iter().map(DuckDbQueryResult::new).collect())
}
#[instrument(skip(self, params), fields(sql = %sql))]
pub async fn raw_sql_execute(&self, sql: &str, params: &[FilterValue]) -> DuckDbResult<u64> {
debug!("Executing raw SQL statement");
let conn = self.connection().await?;
let affected = conn.execute(sql, params).await?;
Ok(affected as u64)
}
#[instrument(skip(self, sql))]
pub async fn raw_sql(&self, sql: prax_query::raw::Sql) -> DuckDbResult<Vec<DuckDbQueryResult>> {
let (query_string, params) = sql.build();
debug!(sql = %query_string, "Executing raw SQL from builder");
self.execute_raw(&query_string, ¶ms).await
}
#[instrument(skip(self, params), fields(sql = %sql))]
pub async fn raw_sql_first(
&self,
sql: &str,
params: &[FilterValue],
) -> DuckDbResult<DuckDbQueryResult> {
let conn = self.connection().await?;
let result = conn.query_one(sql, params).await?;
Ok(DuckDbQueryResult::new(result))
}
#[instrument(skip(self, params), fields(sql = %sql))]
pub async fn raw_sql_optional(
&self,
sql: &str,
params: &[FilterValue],
) -> DuckDbResult<Option<DuckDbQueryResult>> {
let conn = self.connection().await?;
let result = conn.query_optional(sql, params).await?;
Ok(result.map(DuckDbQueryResult::new))
}
#[instrument(skip(self, params), fields(sql = %sql))]
pub async fn raw_sql_scalar<T>(&self, sql: &str, params: &[FilterValue]) -> DuckDbResult<T>
where
T: for<'a> serde::Deserialize<'a>,
{
let conn = self.connection().await?;
let result = conn.query_one(sql, params).await?;
let value = result
.as_object()
.and_then(|obj| obj.values().next())
.ok_or_else(|| DuckDbError::query("raw_sql_scalar returned empty row"))?;
serde_json::from_value(value.clone()).map_err(|e| {
DuckDbError::deserialization(format!("failed to deserialize scalar: {}", e))
})
}
#[instrument(skip(self), fields(sql_len = %sql.len()))]
pub async fn raw_sql_batch(&self, sql: &str) -> DuckDbResult<()> {
debug!("Executing raw SQL batch");
let conn = self.connection().await?;
conn.execute_batch(sql).await
}
#[instrument(skip(self, filters), fields(table = %table))]
pub async fn count(
&self,
table: &str,
filters: &HashMap<String, FilterValue>,
) -> DuckDbResult<u64> {
let mut sql = format!("SELECT COUNT(*) as count FROM \"{}\"", table);
let mut params: Vec<FilterValue> = Vec::new();
if !filters.is_empty() {
let mut conditions = Vec::new();
for (field, value) in filters {
match value {
FilterValue::Null => {
conditions.push(format!("\"{}\" IS NULL", field));
}
_ => {
conditions.push(format!("\"{}\" = ?", field));
params.push(value.clone());
}
}
}
sql.push_str(" WHERE ");
sql.push_str(&conditions.join(" AND "));
}
debug!(sql = %sql, "Executing count");
let conn = self.connection().await?;
let results = conn.query(&sql, ¶ms).await?;
let count = results
.first()
.and_then(|row| row.get("count"))
.and_then(|v| v.as_i64())
.ok_or_else(|| {
DuckDbError::deserialization(
"count query returned no integer 'count' cell".to_string(),
)
})?;
Ok(count as u64)
}
#[instrument(skip(self), fields(query_len = %query.len()))]
pub async fn copy_to_parquet(&self, query: &str, path: &str) -> DuckDbResult<()> {
let conn = self.connection().await?;
conn.copy_to_parquet(query, path).await
}
#[instrument(skip(self), fields(query_len = %query.len()))]
pub async fn copy_to_csv(&self, query: &str, path: &str, header: bool) -> DuckDbResult<()> {
let conn = self.connection().await?;
conn.copy_to_csv(query, path, header).await
}
pub async fn query_parquet(&self, path: &str) -> DuckDbResult<Vec<DuckDbQueryResult>> {
let conn = self.connection().await?;
let results = conn.query_parquet(path).await?;
Ok(results.into_iter().map(DuckDbQueryResult::new).collect())
}
pub async fn query_csv(
&self,
path: &str,
header: bool,
) -> DuckDbResult<Vec<DuckDbQueryResult>> {
let conn = self.connection().await?;
let results = conn.query_csv(path, header).await?;
Ok(results.into_iter().map(DuckDbQueryResult::new).collect())
}
pub async fn query_json(&self, path: &str) -> DuckDbResult<Vec<DuckDbQueryResult>> {
let conn = self.connection().await?;
let results = conn.query_json(path).await?;
Ok(results.into_iter().map(DuckDbQueryResult::new).collect())
}
pub async fn version(&self) -> DuckDbResult<String> {
let result = self.raw_sql_first("SELECT version()", &[]).await?;
result
.data
.as_object()
.and_then(|obj| obj.values().next())
.and_then(|v| v.as_str())
.map(String::from)
.ok_or_else(|| DuckDbError::query("Failed to get version"))
}
pub async fn explain(&self, query: &str) -> DuckDbResult<String> {
let sql = format!("EXPLAIN {}", query);
let results = self.execute_raw(&sql, &[]).await?;
let mut plan = String::new();
for result in results {
if let Some(obj) = result.data.as_object() {
for value in obj.values() {
if let Some(s) = value.as_str() {
plan.push_str(s);
plan.push('\n');
}
}
}
}
Ok(plan)
}
}
fn json_cell_to_filter_value(value: &JsonValue) -> FilterValue {
match value {
JsonValue::Null => FilterValue::Null,
JsonValue::Bool(b) => FilterValue::Bool(*b),
JsonValue::Number(n) => {
if let Some(i) = n.as_i64() {
FilterValue::Int(i)
} else if let Some(f) = n.as_f64() {
FilterValue::Float(f)
} else {
FilterValue::Null
}
}
JsonValue::String(s) => FilterValue::String(s.clone()),
other => FilterValue::Json(other.clone()),
}
}
impl DuckDbEngine {
async fn fetch_typed<T: prax_query::traits::Model + prax_query::row::FromRow>(
&self,
sql: &str,
params: &[FilterValue],
) -> prax_query::QueryResult<Vec<T>> {
let conn = self
.connection()
.await
.map_err(|e| prax_query::QueryError::connection(e.to_string()).with_source(e))?;
let snapshots = conn
.query_rows(sql, params)
.await
.map_err(|e| prax_query::QueryError::database(e.to_string()).with_source(e))?;
snapshots
.into_iter()
.map(|r| {
T::from_row(&r).map_err(|e| {
let msg = e.to_string();
prax_query::QueryError::deserialization(msg).with_source(e)
})
})
.collect()
}
async fn fetch_affected(
&self,
sql: &str,
params: &[FilterValue],
) -> prax_query::QueryResult<u64> {
let conn = self
.connection()
.await
.map_err(|e| prax_query::QueryError::connection(e.to_string()).with_source(e))?;
let affected = conn
.execute(sql, params)
.await
.map_err(|e| prax_query::QueryError::database(e.to_string()).with_source(e))?;
Ok(affected as u64)
}
}
struct RollbackOnPanic {
conn: Arc<PooledConnection>,
armed: bool,
}
impl RollbackOnPanic {
fn new(conn: Arc<PooledConnection>) -> Self {
Self { conn, armed: true }
}
fn disarm(mut self) {
self.armed = false;
}
}
impl Drop for RollbackOnPanic {
fn drop(&mut self) {
if !self.armed {
return;
}
if let Ok(rt) = tokio::runtime::Handle::try_current() {
let conn = self.conn.clone();
rt.spawn(async move {
if let Err(e) = conn.execute_batch("ROLLBACK").await {
tracing::warn!(
error = %e,
"panic-guard ROLLBACK failed; retiring connection"
);
conn.poison();
}
});
} else {
if self.conn.connection().execute_batch("ROLLBACK").is_err() {
self.conn.poison();
}
}
}
}
impl prax_query::traits::QueryEngine for DuckDbEngine {
fn dialect(&self) -> &dyn prax_query::dialect::SqlDialect {
&prax_query::dialect::Postgres
}
fn query_many<T: prax_query::traits::Model + prax_query::row::FromRow + Send + 'static>(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<Vec<T>>> {
let sql = sql.to_string();
Box::pin(async move { self.fetch_typed::<T>(&sql, ¶ms).await })
}
fn query_one<T: prax_query::traits::Model + prax_query::row::FromRow + Send + 'static>(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<T>> {
let sql = sql.to_string();
Box::pin(async move {
let mut rows: Vec<T> = self.fetch_typed::<T>(&sql, ¶ms).await?;
if rows.is_empty() {
Err(prax_query::QueryError::not_found(T::MODEL_NAME))
} else {
Ok(rows.swap_remove(0))
}
})
}
fn query_optional<T: prax_query::traits::Model + prax_query::row::FromRow + Send + 'static>(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<Option<T>>> {
let sql = sql.to_string();
Box::pin(async move {
let mut rows: Vec<T> = self.fetch_typed::<T>(&sql, ¶ms).await?;
Ok(rows.drain(..).next())
})
}
fn execute_insert<T: prax_query::traits::Model + prax_query::row::FromRow + Send + 'static>(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<T>> {
let sql = sql.to_string();
Box::pin(async move {
let mut rows: Vec<T> = self.fetch_typed::<T>(&sql, ¶ms).await?;
if rows.is_empty() {
Err(prax_query::QueryError::deserialization(
"INSERT ... RETURNING produced no row".to_string(),
))
} else {
Ok(rows.swap_remove(0))
}
})
}
fn execute_update<T: prax_query::traits::Model + prax_query::row::FromRow + Send + 'static>(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<Vec<T>>> {
let sql = sql.to_string();
Box::pin(async move { self.fetch_typed::<T>(&sql, ¶ms).await })
}
fn execute_delete(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<u64>> {
let sql = sql.to_string();
Box::pin(async move { self.fetch_affected(&sql, ¶ms).await })
}
fn execute_raw(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<u64>> {
let sql = sql.to_string();
Box::pin(async move { self.fetch_affected(&sql, ¶ms).await })
}
fn count(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<u64>> {
let sql = sql.to_string();
Box::pin(async move {
let conn = self
.connection()
.await
.map_err(|e| prax_query::QueryError::connection(e.to_string()).with_source(e))?;
let snapshots = conn
.query_rows(&sql, ¶ms)
.await
.map_err(|e| prax_query::QueryError::database(e.to_string()).with_source(e))?;
let first = snapshots.into_iter().next().ok_or_else(|| {
prax_query::QueryError::deserialization("count returned no row".to_string())
})?;
use prax_query::row::RowRef;
if let Ok(n) = first.get_i64("count") {
return Ok(n as u64);
}
if let Ok(n) = first.get_i64("count_star()") {
return Ok(n as u64);
}
Err(prax_query::QueryError::deserialization(
"count column missing from DuckDB result".to_string(),
))
})
}
fn aggregate_query(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> prax_query::traits::BoxFuture<'_, prax_query::QueryResult<Vec<HashMap<String, FilterValue>>>>
{
let sql = sql.to_string();
Box::pin(async move {
let conn = self
.connection()
.await
.map_err(|e| prax_query::QueryError::connection(e.to_string()).with_source(e))?;
let rows = conn
.query(&sql, ¶ms)
.await
.map_err(|e| prax_query::QueryError::database(e.to_string()).with_source(e))?;
Ok(rows
.iter()
.map(|row| {
let mut map = HashMap::new();
if let JsonValue::Object(obj) = row {
for (name, value) in obj {
map.insert(name.clone(), json_cell_to_filter_value(value));
}
}
map
})
.collect())
})
}
fn in_transaction(&self) -> bool {
self.tx_conn.is_some()
}
fn transaction<'a, R, Fut, F>(
&'a self,
f: F,
) -> prax_query::traits::BoxFuture<'a, prax_query::QueryResult<R>>
where
F: FnOnce(Self) -> Fut + Send + 'a,
Fut: std::future::Future<Output = prax_query::QueryResult<R>> + Send + 'a,
R: Send + 'a,
Self: Clone,
{
Box::pin(async move {
if self.tx_conn.is_some() {
return Err(prax_query::QueryError::internal(
"nested transactions not yet implemented \
(call .transaction() on the outer engine only, or \
issue SAVEPOINT via execute_raw)",
));
}
let conn =
self.pool.get().await.map_err(|e| {
prax_query::QueryError::connection(e.to_string()).with_source(e)
})?;
conn.execute_batch("BEGIN TRANSACTION")
.await
.map_err(|e| prax_query::QueryError::database(e.to_string()).with_source(e))?;
let tx_conn = Arc::new(conn);
let finalized = Arc::new(AtomicBool::new(false));
let tx_engine = DuckDbEngine {
pool: self.pool.clone(),
tx_conn: Some(tx_conn.clone()),
tx_finalized: Some(finalized.clone()),
};
let rollback_guard = RollbackOnPanic::new(tx_conn.clone());
let result = f(tx_engine).await;
finalized.store(true, Ordering::Release);
match result {
Ok(v) => match tx_conn.execute_batch("COMMIT").await {
Ok(()) => {
rollback_guard.disarm();
Ok(v)
}
Err(e) => {
if let Err(rb) = tx_conn.execute_batch("ROLLBACK").await {
tracing::warn!(
error = %rb,
"ROLLBACK after failed COMMIT failed; retiring connection"
);
tx_conn.poison();
}
rollback_guard.disarm();
Err(prax_query::QueryError::database(e.to_string()).with_source(e))
}
},
Err(e) => {
if let Err(rb) = tx_conn.execute_batch("ROLLBACK").await {
tracing::warn!(
error = %rb,
"ROLLBACK failed; retiring connection instead of returning it to the idle pool"
);
tx_conn.poison();
}
rollback_guard.disarm();
Err(e)
}
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::DuckDbConfig;
#[tokio::test]
async fn test_engine_creation() {
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let engine = DuckDbEngine::new(pool);
let version = engine.version().await.unwrap();
assert!(!version.is_empty());
}
#[tokio::test]
async fn test_query_many() {
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let engine = DuckDbEngine::new(pool);
engine
.raw_sql_batch(
"CREATE TABLE test (id INTEGER, name VARCHAR);
INSERT INTO test VALUES (1, 'Alice'), (2, 'Bob');",
)
.await
.unwrap();
let results = engine
.query_many("test", &[], &HashMap::new(), &[], None, None)
.await
.unwrap();
assert_eq!(results.len(), 2);
}
#[tokio::test]
async fn test_count() {
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let engine = DuckDbEngine::new(pool);
engine
.raw_sql_batch(
"CREATE TABLE test (id INTEGER);
INSERT INTO test VALUES (1), (2), (3);",
)
.await
.unwrap();
let count = engine.count("test", &HashMap::new()).await.unwrap();
assert_eq!(count, 3);
}
#[tokio::test]
async fn test_transaction_rollback() {
use prax_query::traits::QueryEngine;
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let engine = DuckDbEngine::new(pool);
engine
.raw_sql_batch("CREATE TABLE tx_test (id INTEGER);")
.await
.unwrap();
let result: prax_query::QueryResult<()> = engine
.transaction(|tx| async move {
assert!(tx.in_transaction());
QueryEngine::execute_raw(&tx, "INSERT INTO tx_test VALUES (1)", vec![]).await?;
Err(prax_query::QueryError::internal("forced failure"))
})
.await;
assert!(result.is_err());
assert!(!engine.in_transaction());
let count = engine.count("tx_test", &HashMap::new()).await.unwrap();
assert_eq!(count, 0);
}
#[tokio::test]
async fn test_transaction_commit() {
use prax_query::traits::QueryEngine;
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let engine = DuckDbEngine::new(pool);
engine
.raw_sql_batch("CREATE TABLE tx_commit (id INTEGER);")
.await
.unwrap();
let result: prax_query::QueryResult<()> = engine
.transaction(|tx| async move {
QueryEngine::execute_raw(&tx, "INSERT INTO tx_commit VALUES (1)", vec![]).await?;
Ok(())
})
.await;
assert!(result.is_ok());
let count = engine.count("tx_commit", &HashMap::new()).await.unwrap();
assert_eq!(count, 1);
}
#[tokio::test]
async fn test_aggregate_query() {
use prax_query::traits::QueryEngine;
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let engine = DuckDbEngine::new(pool);
engine
.raw_sql_batch(
"CREATE TABLE agg (id INTEGER, grp VARCHAR);
INSERT INTO agg VALUES (1, 'a'), (2, 'a'), (3, 'b');",
)
.await
.unwrap();
let rows = QueryEngine::aggregate_query(
&engine,
"SELECT grp, COUNT(*) AS cnt, AVG(id) AS avg_id FROM agg GROUP BY grp ORDER BY grp",
vec![],
)
.await
.unwrap();
assert_eq!(rows.len(), 2);
assert_eq!(
rows[0].get("grp"),
Some(&FilterValue::String("a".to_string()))
);
assert_eq!(rows[0].get("cnt"), Some(&FilterValue::Int(2)));
assert_eq!(rows[0].get("avg_id"), Some(&FilterValue::Float(1.5)));
assert_eq!(
rows[1].get("grp"),
Some(&FilterValue::String("b".to_string()))
);
assert_eq!(rows[1].get("cnt"), Some(&FilterValue::Int(1)));
assert_eq!(rows[1].get("avg_id"), Some(&FilterValue::Float(3.0)));
}
#[tokio::test]
async fn test_stashed_tx_engine_fails_after_finalize() {
use prax_query::traits::QueryEngine;
let pool = DuckDbPool::new(DuckDbConfig::in_memory()).await.unwrap();
let engine = DuckDbEngine::new(pool);
engine
.raw_sql_batch("CREATE TABLE stash (id INTEGER);")
.await
.unwrap();
let tx: DuckDbEngine = engine
.transaction(|tx| async move { Ok::<_, prax_query::QueryError>(tx) })
.await
.unwrap();
let err = QueryEngine::execute_raw(&tx, "INSERT INTO stash VALUES (1)", vec![])
.await
.unwrap_err();
assert!(
err.to_string().contains(TX_FINALIZED),
"expected a finalized-transaction error, got: {err}"
);
drop(tx);
let count = engine.count("stash", &HashMap::new()).await.unwrap();
assert_eq!(count, 0);
}
}