use arrow::{datatypes::Schema, record_batch::RecordBatch};
use duckdb::{Connection, Result as DuckDBResult, vtab::arrow_recordbatch_to_query_params};
use std::sync::Arc;
fn table_exists(conn: &Connection, table_name: &str) -> DuckDBResult<bool> {
let exists: bool = conn.query_row(
"SELECT EXISTS (SELECT 1 FROM information_schema.tables WHERE table_name = ?)",
[table_name],
|row| row.get(0),
)?;
Ok(exists)
}
pub fn save_to_duckdb(
conn: &Connection,
table_name: &str,
_schema: Arc<Schema>,
batches: Vec<RecordBatch>,
) -> DuckDBResult<()> {
let mut batches = batches.into_iter();
if !table_exists(conn, table_name)?
&& let Some(first) = batches.next()
{
conn.execute(
&format!("CREATE TABLE {table_name} AS SELECT * FROM arrow(?, ?)"),
arrow_recordbatch_to_query_params(first),
)?;
}
for batch in batches {
conn.execute(
&format!("INSERT INTO {table_name} SELECT * FROM arrow(?, ?)"),
arrow_recordbatch_to_query_params(batch),
)?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::{Float64Array, Int64Array, StringArray};
use arrow::datatypes::{DataType, Field, Schema};
use duckdb::vtab::arrow::ArrowVTab;
fn duck_db_conn() -> Connection {
let conn = Connection::open_in_memory().unwrap();
conn.register_table_function::<ArrowVTab>("arrow").unwrap();
conn
}
fn sample_batch() -> (Arc<Schema>, RecordBatch) {
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int64, false),
Field::new("name", DataType::Utf8, false),
Field::new("score", DataType::Float64, true),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int64Array::from(vec![1, 2, 3])),
Arc::new(StringArray::from(vec!["a", "b", "c"])),
Arc::new(Float64Array::from(vec![Some(1.1), None, Some(3.3)])),
],
)
.unwrap();
(schema, batch)
}
#[test]
fn create_table_with_schema_and_data() {
let conn = duck_db_conn();
let (schema, batch) = sample_batch();
save_to_duckdb(&conn, "my_table", schema, vec![batch]).unwrap();
let count: i64 = conn
.query_row("SELECT COUNT(*) FROM my_table", [], |row| row.get(0))
.unwrap();
assert_eq!(count, 3);
let name: String = conn
.query_row("SELECT name FROM my_table WHERE id = 1", [], |row| {
row.get(0)
})
.unwrap();
assert_eq!(name, "a");
}
#[test]
fn append_multiple_batches() {
let conn = duck_db_conn();
let (schema, batch1) = sample_batch();
let (_, batch2) = sample_batch();
save_to_duckdb(&conn, "my_table", schema, vec![batch1, batch2]).unwrap();
let count: i64 = conn
.query_row("SELECT COUNT(*) FROM my_table", [], |row| row.get(0))
.unwrap();
assert_eq!(count, 6);
}
#[test]
fn no_batches_does_not_create_table() {
let conn = duck_db_conn();
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int64, false)]));
let result = save_to_duckdb(&conn, "my_table", schema, vec![]);
assert!(result.is_ok());
let table_exists = conn
.query_row(
"SELECT COUNT(*) FROM information_schema.tables WHERE table_name = 'my_table'",
[],
|row| row.get::<_, i64>(0),
)
.unwrap();
assert_eq!(
table_exists, 0,
"no table should exist when batches is empty"
);
}
}