diff --git a/include/lbug_arrow.h b/include/lbug_arrow.h index eb2ec23..6c6b358 100644 --- a/include/lbug_arrow.h +++ b/include/lbug_arrow.h @@ -10,6 +10,7 @@ namespace lbug_arrow { ArrowSchema query_result_get_arrow_schema(const lbug::main::QueryResult& result); +ArrowSchema prepared_statement_get_arrow_schema(const lbug::main::PreparedStatement& stmt); bool query_result_has_next_arrow_chunk(lbug::main::QueryResult& result); ArrowArray query_result_get_next_arrow_chunk(lbug::main::QueryResult& result, uint64_t chunkSize); ArrowArray query_result_get_csr_indptr(const lbug::main::QueryResult& result); diff --git a/src/connection.rs b/src/connection.rs index 5a5a6bf..1910274 100644 --- a/src/connection.rs +++ b/src/connection.rs @@ -56,6 +56,19 @@ impl PreparedStatement { pub fn is_read_only(&self) -> bool { ffi::prepared_statement_is_read_only(&self.statement) } + + #[cfg(feature = "arrow")] + /// Returns the query's result schema as Arrow without executing it. + /// + /// The schema comes from the bound and planned statement, so this is + /// cheap even for expensive queries. + /// + /// *Requires the `arrow` feature* + pub fn get_arrow_schema(&self) -> Result { + Ok( + crate::ffi::arrow::ffi_arrow::prepared_statement_get_arrow_schema(&self.statement)?.0, + ) + } } /// Connections are used to interact with a Database instance. diff --git a/src/ffi/arrow.rs b/src/ffi/arrow.rs index 564b8ac..27e8dc5 100644 --- a/src/ffi/arrow.rs +++ b/src/ffi/arrow.rs @@ -29,6 +29,9 @@ pub(crate) mod ffi_arrow { #[namespace = "lbug::main"] type QueryResult<'db> = crate::ffi::ffi::QueryResult<'db>; + + #[namespace = "lbug::main"] + type PreparedStatement = crate::ffi::ffi::PreparedStatement; } unsafe extern "C++" { @@ -61,6 +64,11 @@ pub(crate) mod ffi_arrow { #[namespace = "lbug_arrow"] fn query_result_get_arrow_schema<'db>(result: &QueryResult<'db>) -> Result; + + #[namespace = "lbug_arrow"] + fn prepared_statement_get_arrow_schema( + statement: &PreparedStatement, + ) -> Result; } #[namespace = "lbug_rs"] diff --git a/src/lbug_arrow.cpp b/src/lbug_arrow.cpp index a5ac1c1..fbc7f5a 100644 --- a/src/lbug_arrow.cpp +++ b/src/lbug_arrow.cpp @@ -4,6 +4,8 @@ #include #include +#include "common/arrow/arrow_converter.h" + namespace lbug { namespace main { @@ -69,6 +71,12 @@ ArrowSchema query_result_get_arrow_schema(const lbug::main::QueryResult& result) return *result.getArrowSchema(); } +ArrowSchema prepared_statement_get_arrow_schema(const lbug::main::PreparedStatement& stmt) { + // Schema only: names/types come from the bound and planned statement, no execution. + return *lbug::common::ArrowConverter::toArrowSchema(stmt.getColumnTypes(), + stmt.getColumnNames(), false /* fallbackExtensionTypes */); +} + bool query_result_has_next_arrow_chunk(lbug::main::QueryResult& result) { return result.hasNextArrowChunk(); }