use std::collections::HashMap;
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;
use crate::types::filter_value_to_json;
#[derive(Clone)]
pub struct DuckDbEngine {
pool: DuckDbPool,
}
#[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 }
}
pub fn pool(&self) -> &DuckDbPool {
&self.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.pool.get().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.pool.get().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.pool.get().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);
debug!(sql = %sql, "Executing insert");
let conn = self.pool.get().await?;
conn.execute(&sql, ¶ms).await?;
let json = data
.iter()
.map(|(k, v)| (k.clone(), filter_value_to_json(v)))
.collect::<serde_json::Map<_, _>>();
Ok(DuckDbQueryResult::new(JsonValue::Object(json)))
}
#[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.pool.get().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.pool.get().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.pool.get().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.pool.get().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.pool.get().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.pool.get().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.pool.get().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.pool.get().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.pool.get().await?;
let results = conn.query(&sql, ¶ms).await?;
let count = results
.first()
.and_then(|row| row.get("count"))
.and_then(|v| v.as_i64())
.unwrap_or(0);
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.pool.get().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.pool.get().await?;
conn.copy_to_csv(query, path, header).await
}
pub async fn query_parquet(&self, path: &str) -> DuckDbResult<Vec<DuckDbQueryResult>> {
let conn = self.pool.get().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.pool.get().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.pool.get().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)
}
}
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
.pool
.get()
.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
.pool
.get()
.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)
}
}
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.pool.get().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(),
))
})
}
}
#[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);
}
}