use std::ffi::{CStr, CString};
use libduckdb_sys::{
duckdb_bind_blob, duckdb_bind_boolean, duckdb_bind_double, duckdb_bind_int64, duckdb_bind_null,
duckdb_bind_parameter_index, duckdb_bind_varchar_length, duckdb_clear_bindings,
duckdb_column_count, duckdb_column_name, duckdb_column_type, duckdb_connect, duckdb_connection,
duckdb_database, duckdb_destroy_data_chunk, duckdb_destroy_prepare, duckdb_destroy_result,
duckdb_disconnect, duckdb_execute_prepared, duckdb_fetch_chunk, duckdb_nparams,
duckdb_parameter_name, duckdb_prepare, duckdb_prepare_error, duckdb_prepared_statement,
duckdb_query, duckdb_result, duckdb_result_error, duckdb_rows_changed, idx_t, DuckDBSuccess,
};
use crate::data_chunk::DataChunk;
use crate::error::ExtensionError;
use crate::types::TypeId;
fn to_c_sql(sql: &str) -> Result<CString, ExtensionError> {
CString::new(sql)
.map_err(|_| ExtensionError::new("SQL text must not contain an interior NUL byte"))
}
unsafe fn c_str_to_owned(ptr: *const std::os::raw::c_char) -> Option<String> {
if ptr.is_null() {
return None;
}
unsafe { CStr::from_ptr(ptr) }
.to_str()
.ok()
.map(str::to_owned)
}
pub struct OwnedDataChunk {
chunk: libduckdb_sys::duckdb_data_chunk,
view: DataChunk,
}
impl OwnedDataChunk {
#[must_use]
pub const unsafe fn from_raw(chunk: libduckdb_sys::duckdb_data_chunk) -> Self {
Self {
chunk,
view: unsafe { DataChunk::from_raw(chunk) },
}
}
#[must_use]
pub const fn into_raw(self) -> libduckdb_sys::duckdb_data_chunk {
let raw = self.chunk;
std::mem::forget(self);
raw
}
}
impl std::fmt::Debug for OwnedDataChunk {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OwnedDataChunk")
.field("chunk", &self.chunk)
.finish_non_exhaustive()
}
}
impl std::ops::Deref for OwnedDataChunk {
type Target = DataChunk;
fn deref(&self) -> &DataChunk {
&self.view
}
}
impl Drop for OwnedDataChunk {
fn drop(&mut self) {
unsafe { duckdb_destroy_data_chunk(&raw mut self.chunk) };
}
}
pub struct QueryResult {
result: duckdb_result,
}
impl QueryResult {
#[must_use]
pub fn column_count(&self) -> usize {
let mut result = self.result;
usize::try_from(unsafe { duckdb_column_count(&raw mut result) }).unwrap_or(0)
}
#[must_use]
pub fn column_name(&self, index: usize) -> Option<String> {
if index >= self.column_count() {
return None;
}
let mut result = self.result;
let ptr = unsafe { duckdb_column_name(&raw mut result, index as idx_t) };
unsafe { c_str_to_owned(ptr) }
}
#[must_use]
pub fn column_type(&self, index: usize) -> Option<TypeId> {
if index >= self.column_count() {
return None;
}
let mut result = self.result;
let raw = unsafe { duckdb_column_type(&raw mut result, index as idx_t) };
TypeId::try_from_duckdb_type(raw)
}
#[must_use]
pub fn rows_changed(&self) -> u64 {
let mut result = self.result;
unsafe { duckdb_rows_changed(&raw mut result) }
}
#[must_use]
pub fn next_chunk(&mut self) -> Option<OwnedDataChunk> {
let chunk = unsafe { duckdb_fetch_chunk(self.result) };
if chunk.is_null() {
return None;
}
Some(unsafe { OwnedDataChunk::from_raw(chunk) })
}
#[must_use]
pub const fn as_raw(&self) -> &duckdb_result {
&self.result
}
}
impl std::fmt::Debug for QueryResult {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("QueryResult").finish_non_exhaustive()
}
}
impl Drop for QueryResult {
fn drop(&mut self) {
unsafe { duckdb_destroy_result(&raw mut self.result) };
}
}
pub unsafe fn query(con: duckdb_connection, sql: &str) -> Result<QueryResult, ExtensionError> {
let c_sql = to_c_sql(sql)?;
let mut result: duckdb_result = unsafe { std::mem::zeroed() };
let state = unsafe { duckdb_query(con, c_sql.as_ptr(), &raw mut result) };
if state == DuckDBSuccess {
return Ok(QueryResult { result });
}
let message = unsafe { c_str_to_owned(duckdb_result_error(&raw mut result)) }
.unwrap_or_else(|| String::from("query failed without an error message"));
unsafe { duckdb_destroy_result(&raw mut result) };
Err(ExtensionError::new(message))
}
pub unsafe fn execute(con: duckdb_connection, sql: &str) -> Result<u64, ExtensionError> {
let result = unsafe { query(con, sql) }?;
Ok(result.rows_changed())
}
pub struct PreparedStatement {
statement: duckdb_prepared_statement,
}
impl PreparedStatement {
#[must_use]
pub fn parameter_count(&self) -> usize {
usize::try_from(unsafe { duckdb_nparams(self.statement) }).unwrap_or(0)
}
#[must_use]
pub fn parameter_name(&self, index: usize) -> Option<String> {
if index == 0 || index > self.parameter_count() {
return None;
}
let ptr = unsafe { duckdb_parameter_name(self.statement, index as idx_t) };
let name = unsafe { c_str_to_owned(ptr) };
if !ptr.is_null() {
unsafe { libduckdb_sys::duckdb_free(ptr.cast_mut().cast()) };
}
name
}
#[must_use]
pub fn parameter_index(&self, name: &str) -> Option<usize> {
let c_name = CString::new(name).ok()?;
let mut index: idx_t = 0;
let state =
unsafe { duckdb_bind_parameter_index(self.statement, &raw mut index, c_name.as_ptr()) };
(state == DuckDBSuccess).then(|| usize::try_from(index).unwrap_or(0))
}
pub fn bind_i64(&self, index: usize, value: i64) -> Result<(), ExtensionError> {
self.check(
unsafe { duckdb_bind_int64(self.statement, index as idx_t, value) },
index,
)
}
pub fn bind_f64(&self, index: usize, value: f64) -> Result<(), ExtensionError> {
self.check(
unsafe { duckdb_bind_double(self.statement, index as idx_t, value) },
index,
)
}
pub fn bind_bool(&self, index: usize, value: bool) -> Result<(), ExtensionError> {
self.check(
unsafe { duckdb_bind_boolean(self.statement, index as idx_t, value) },
index,
)
}
pub fn bind_str(&self, index: usize, value: &str) -> Result<(), ExtensionError> {
let state = unsafe {
duckdb_bind_varchar_length(
self.statement,
index as idx_t,
value.as_ptr().cast::<std::os::raw::c_char>(),
idx_t::try_from(value.len()).unwrap_or(idx_t::MAX),
)
};
self.check(state, index)
}
pub fn bind_blob(&self, index: usize, value: &[u8]) -> Result<(), ExtensionError> {
let state = unsafe {
duckdb_bind_blob(
self.statement,
index as idx_t,
value.as_ptr().cast::<std::os::raw::c_void>(),
idx_t::try_from(value.len()).unwrap_or(idx_t::MAX),
)
};
self.check(state, index)
}
pub fn bind_null(&self, index: usize) -> Result<(), ExtensionError> {
self.check(
unsafe { duckdb_bind_null(self.statement, index as idx_t) },
index,
)
}
pub fn clear_bindings(&self) -> Result<(), ExtensionError> {
if unsafe { duckdb_clear_bindings(self.statement) } == DuckDBSuccess {
Ok(())
} else {
Err(ExtensionError::new("duckdb_clear_bindings failed"))
}
}
pub fn execute(&self) -> Result<QueryResult, ExtensionError> {
let mut result: duckdb_result = unsafe { std::mem::zeroed() };
let state = unsafe { duckdb_execute_prepared(self.statement, &raw mut result) };
if state == DuckDBSuccess {
return Ok(QueryResult { result });
}
let message = unsafe { c_str_to_owned(duckdb_result_error(&raw mut result)) }
.unwrap_or_else(|| String::from("prepared statement failed without an error message"));
unsafe { duckdb_destroy_result(&raw mut result) };
Err(ExtensionError::new(message))
}
#[must_use]
pub const fn as_raw(&self) -> duckdb_prepared_statement {
self.statement
}
fn check(
&self,
state: libduckdb_sys::duckdb_state,
index: usize,
) -> Result<(), ExtensionError> {
if state == DuckDBSuccess {
Ok(())
} else {
Err(ExtensionError::new(format!(
"failed to bind parameter {index} (statement has {} parameter(s))",
self.parameter_count()
)))
}
}
}
impl std::fmt::Debug for PreparedStatement {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PreparedStatement")
.field("statement", &self.statement)
.finish()
}
}
impl Drop for PreparedStatement {
fn drop(&mut self) {
unsafe { duckdb_destroy_prepare(&raw mut self.statement) };
}
}
pub unsafe fn prepare(
con: duckdb_connection,
sql: &str,
) -> Result<PreparedStatement, ExtensionError> {
let c_sql = to_c_sql(sql)?;
let mut statement: duckdb_prepared_statement = std::ptr::null_mut();
let state = unsafe { duckdb_prepare(con, c_sql.as_ptr(), &raw mut statement) };
if state == DuckDBSuccess {
return Ok(PreparedStatement { statement });
}
let message = unsafe { c_str_to_owned(duckdb_prepare_error(statement)) }
.unwrap_or_else(|| String::from("prepare failed without an error message"));
unsafe { duckdb_destroy_prepare(&raw mut statement) };
Err(ExtensionError::new(message))
}
pub struct OwnedConnection {
con: duckdb_connection,
}
impl OwnedConnection {
pub unsafe fn open(db: duckdb_database) -> Result<Self, ExtensionError> {
let mut con: duckdb_connection = std::ptr::null_mut();
if unsafe { duckdb_connect(db, &raw mut con) } != DuckDBSuccess {
return Err(ExtensionError::new("duckdb_connect failed"));
}
Ok(Self { con })
}
pub fn query(&self, sql: &str) -> Result<QueryResult, ExtensionError> {
unsafe { query(self.con, sql) }
}
pub fn execute(&self, sql: &str) -> Result<u64, ExtensionError> {
unsafe { execute(self.con, sql) }
}
pub fn prepare(&self, sql: &str) -> Result<PreparedStatement, ExtensionError> {
unsafe { prepare(self.con, sql) }
}
#[must_use]
pub const fn as_raw(&self) -> duckdb_connection {
self.con
}
}
impl std::fmt::Debug for OwnedConnection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OwnedConnection")
.field("con", &self.con)
.finish()
}
}
impl Drop for OwnedConnection {
fn drop(&mut self) {
unsafe { duckdb_disconnect(&raw mut self.con) };
}
}
unsafe impl Send for OwnedConnection {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_sql_containing_an_interior_nul() {
let err = to_c_sql("SELECT 1\0; DROP TABLE t").expect_err("must reject NUL");
assert!(err.as_str().contains("NUL"), "{err}");
}
#[test]
fn accepts_ordinary_sql() {
assert_eq!(
to_c_sql("SELECT 1")
.expect("valid SQL")
.to_str()
.expect("utf8"),
"SELECT 1"
);
}
#[test]
fn null_c_string_reads_as_none() {
assert_eq!(unsafe { c_str_to_owned(std::ptr::null()) }, None);
}
#[test]
fn owned_connection_is_send() {
const fn assert_send<T: Send>() {}
assert_send::<OwnedConnection>();
}
}
#[cfg(all(test, feature = "_duckdb-testing"))]
mod live_tests {
use super::*;
use crate::testing::InMemoryDb;
unsafe fn open_raw() -> (libduckdb_sys::duckdb_database, duckdb_connection) {
let mut db: libduckdb_sys::duckdb_database = std::ptr::null_mut();
let mut con: duckdb_connection = std::ptr::null_mut();
unsafe {
assert_eq!(
libduckdb_sys::duckdb_open(std::ptr::null(), &raw mut db),
DuckDBSuccess
);
assert_eq!(
libduckdb_sys::duckdb_connect(db, &raw mut con),
DuckDBSuccess
);
}
(db, con)
}
unsafe fn close_raw(mut db: libduckdb_sys::duckdb_database, mut con: duckdb_connection) {
unsafe {
libduckdb_sys::duckdb_disconnect(&raw mut con);
libduckdb_sys::duckdb_close(&raw mut db);
}
}
#[test]
fn query_reads_scalar_results() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
let mut result =
unsafe { query(con, "SELECT 42 AS answer, 'hi' AS greeting") }.expect("query succeeds");
assert_eq!(result.column_count(), 2);
assert_eq!(result.column_name(0).as_deref(), Some("answer"));
assert_eq!(result.column_name(1).as_deref(), Some("greeting"));
assert_eq!(result.column_type(0), Some(TypeId::Integer));
assert_eq!(result.column_type(1), Some(TypeId::Varchar));
assert_eq!(result.column_name(2), None);
assert_eq!(result.column_type(2), None);
let chunk = result.next_chunk().expect("one chunk");
assert_eq!(chunk.size(), 1);
unsafe {
assert_eq!(chunk.reader(0).read_i32(0), 42);
assert_eq!(chunk.reader(1).read_str(0), "hi");
}
drop(chunk);
assert!(result.next_chunk().is_none());
drop(result);
unsafe { close_raw(db, con) };
}
#[test]
fn query_surfaces_duckdb_errors() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
let err = unsafe { query(con, "SELECT * FROM no_such_table") }
.expect_err("missing table must fail");
assert!(err.as_str().contains("no_such_table"), "{err}");
let ok = unsafe { query(con, "SELECT 1") };
assert!(ok.is_ok());
unsafe { close_raw(db, con) };
}
#[test]
fn execute_reports_rows_changed() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
unsafe { execute(con, "CREATE TABLE t(i INTEGER)") }.expect("create");
let changed =
unsafe { execute(con, "INSERT INTO t VALUES (1), (2), (3)") }.expect("insert");
assert_eq!(changed, 3);
unsafe { close_raw(db, con) };
}
#[test]
fn multi_chunk_results_are_fully_drained() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
let rows = crate::vector::vector_size() * 3 + 7;
let sql = format!("SELECT i FROM range({rows}) t(i)");
let mut result = unsafe { query(con, &sql) }.expect("range query");
let mut seen: u64 = 0;
let mut chunks = 0;
while let Some(chunk) = result.next_chunk() {
chunks += 1;
for row in 0..chunk.size() {
assert_eq!(
unsafe { chunk.reader(0).read_i64(row) },
i64::try_from(seen).expect("row counter fits in i64")
);
seen += 1;
}
}
assert_eq!(seen, rows);
assert!(chunks > 1, "expected several chunks, got {chunks}");
drop(result);
unsafe { close_raw(db, con) };
}
#[test]
fn prepared_statements_bind_and_execute() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
let stmt = unsafe { prepare(con, "SELECT ? + ?") }.expect("prepare");
assert_eq!(stmt.parameter_count(), 2);
stmt.bind_i64(1, 20).expect("bind 1");
stmt.bind_i64(2, 22).expect("bind 2");
let mut result = stmt.execute().expect("execute");
let chunk = result.next_chunk().expect("one chunk");
assert_eq!(unsafe { chunk.reader(0).read_i64(0) }, 42);
drop(chunk);
drop(result);
drop(stmt);
unsafe { close_raw(db, con) };
}
#[test]
fn prepared_statements_are_reusable_after_clearing() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
let stmt = unsafe { prepare(con, "SELECT ?::BIGINT * 2") }.expect("prepare");
for input in [1_i64, 7, 100] {
stmt.clear_bindings().expect("clear");
stmt.bind_i64(1, input).expect("bind");
let mut result = stmt.execute().expect("execute");
let chunk = result.next_chunk().expect("one chunk");
assert_eq!(unsafe { chunk.reader(0).read_i64(0) }, input * 2);
}
drop(stmt);
unsafe { close_raw(db, con) };
}
#[test]
fn binding_a_string_avoids_sql_injection() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
unsafe { execute(con, "CREATE TABLE t(s VARCHAR)") }.expect("create");
let stmt = unsafe { prepare(con, "INSERT INTO t VALUES (?)") }.expect("prepare");
stmt.bind_str(1, "'); DROP TABLE t; --").expect("bind");
stmt.execute().expect("insert");
drop(stmt);
let mut result = unsafe { query(con, "SELECT s FROM t") }.expect("select");
let chunk = result.next_chunk().expect("one chunk");
assert_eq!(chunk.size(), 1);
assert_eq!(
unsafe { chunk.reader(0).read_str(0) },
"'); DROP TABLE t; --"
);
drop(chunk);
drop(result);
unsafe { close_raw(db, con) };
}
#[test]
fn named_parameters_resolve_by_name() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
let stmt = unsafe { prepare(con, "SELECT $needle::BIGINT") }.expect("prepare");
let index = stmt.parameter_index("needle").expect("named parameter");
assert_eq!(stmt.parameter_name(index).as_deref(), Some("needle"));
assert_eq!(stmt.parameter_index("nope"), None);
stmt.bind_i64(index, 5).expect("bind");
let mut result = stmt.execute().expect("execute");
let chunk = result.next_chunk().expect("one chunk");
assert_eq!(unsafe { chunk.reader(0).read_i64(0) }, 5);
drop(chunk);
drop(result);
drop(stmt);
unsafe { close_raw(db, con) };
}
#[test]
fn prepare_surfaces_parse_errors() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
let err = unsafe { prepare(con, "SELECT FROM WHERE") }.expect_err("syntax error");
assert!(!err.as_str().is_empty());
unsafe { close_raw(db, con) };
}
#[test]
fn null_and_blob_bindings_round_trip() {
let _guard = InMemoryDb::open().expect("dispatch table");
let (db, con) = unsafe { open_raw() };
unsafe { execute(con, "CREATE TABLE t(b BLOB, n INTEGER)") }.expect("create");
let stmt = unsafe { prepare(con, "INSERT INTO t VALUES (?, ?)") }.expect("prepare");
stmt.bind_blob(1, &[0x00, 0xFF, 0x80]).expect("bind blob");
stmt.bind_null(2).expect("bind null");
stmt.execute().expect("insert");
drop(stmt);
let mut result = unsafe { query(con, "SELECT b, n FROM t") }.expect("select");
let chunk = result.next_chunk().expect("one chunk");
unsafe {
assert_eq!(chunk.reader(0).read_blob(0), &[0x00, 0xFF, 0x80]);
assert!(!chunk.reader(1).is_valid(0));
}
drop(chunk);
drop(result);
unsafe { close_raw(db, con) };
}
#[test]
fn owned_connection_outlives_the_database_handle() {
let _guard = InMemoryDb::open().expect("dispatch table");
let mut db: libduckdb_sys::duckdb_database = std::ptr::null_mut();
unsafe {
assert_eq!(
libduckdb_sys::duckdb_open(std::ptr::null(), &raw mut db),
DuckDBSuccess
);
}
let con = unsafe { OwnedConnection::open(db) }.expect("connect");
unsafe { libduckdb_sys::duckdb_close(&raw mut db) };
con.execute("CREATE TABLE t(i INTEGER)").expect("create");
con.execute("INSERT INTO t VALUES (1), (2)")
.expect("insert");
let mut result = con.query("SELECT count(*) FROM t").expect("count");
let chunk = result.next_chunk().expect("one chunk");
assert_eq!(unsafe { chunk.reader(0).read_i64(0) }, 2);
}
}