use libduckdb_sys::{
duckdb_connect, duckdb_connection, duckdb_disconnect, duckdb_extension_access,
duckdb_extension_info, duckdb_rs_extension_api_init, DuckDBSuccess,
};
use crate::abi::AbiPolicy;
use crate::connection::Connection;
use crate::error::ExtensionError;
#[macro_export]
macro_rules! entry_point {
($fn_name:ident, $register:expr) => {
$crate::entry_point!($fn_name, $crate::abi::AbiPolicy::Strict, $register);
};
($fn_name:ident, $policy:expr, $register:expr) => {
#[no_mangle]
pub unsafe extern "C" fn $fn_name(
info: ::libduckdb_sys::duckdb_extension_info,
access: *const ::libduckdb_sys::duckdb_extension_access,
) -> bool {
unsafe {
$crate::entry_point::init_extension_with_policy(
info,
access,
$crate::DUCKDB_API_VERSION,
$policy,
$register,
)
}
}
};
}
#[macro_export]
macro_rules! entry_point_v2 {
($fn_name:ident, $register:expr) => {
$crate::entry_point_v2!($fn_name, $crate::abi::AbiPolicy::Strict, $register);
};
($fn_name:ident, $policy:expr, $register:expr) => {
#[no_mangle]
pub unsafe extern "C" fn $fn_name(
info: ::libduckdb_sys::duckdb_extension_info,
access: *const ::libduckdb_sys::duckdb_extension_access,
) -> bool {
unsafe {
$crate::entry_point::init_extension_v2_with_policy(
info,
access,
$crate::DUCKDB_API_VERSION,
$policy,
$register,
)
}
}
};
}
pub unsafe fn init_extension<F>(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
api_version: &str,
register: F,
) -> bool
where
F: FnOnce(duckdb_connection) -> Result<(), ExtensionError>,
{
unsafe { init_extension_with_policy(info, access, api_version, AbiPolicy::default(), register) }
}
pub unsafe fn init_extension_with_policy<F>(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
api_version: &str,
policy: AbiPolicy,
register: F,
) -> bool
where
F: FnOnce(duckdb_connection) -> Result<(), ExtensionError>,
{
match unsafe { init_extension_internal(info, access, api_version, policy, register) } {
Ok(result) => result,
Err(e) => {
unsafe { report_error(info, access, &e) };
false
}
}
}
pub unsafe fn init_extension_v2<F>(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
api_version: &str,
register: F,
) -> bool
where
F: FnOnce(&Connection) -> Result<(), crate::error::ExtensionError>,
{
unsafe {
init_extension_v2_with_policy(info, access, api_version, AbiPolicy::default(), register)
}
}
pub unsafe fn init_extension_v2_with_policy<F>(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
api_version: &str,
policy: AbiPolicy,
register: F,
) -> bool
where
F: FnOnce(&Connection) -> Result<(), crate::error::ExtensionError>,
{
match unsafe { init_extension_v2_internal(info, access, api_version, policy, register) } {
Ok(result) => result,
Err(e) => {
unsafe { report_error(info, access, &e) };
false
}
}
}
unsafe fn init_extension_v2_internal<F>(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
api_version: &str,
policy: AbiPolicy,
register: F,
) -> Result<bool, crate::error::ExtensionError>
where
F: FnOnce(&Connection) -> Result<(), crate::error::ExtensionError>,
{
if api_version.contains('\0') {
return Err(crate::error::ExtensionError::new(
"api_version must not contain an interior NUL byte",
));
}
let have_api = unsafe {
duckdb_rs_extension_api_init(info, access, api_version)
.map_err(|e| crate::error::ExtensionError::new(e.to_string()))?
};
if !have_api {
return Ok(false);
}
unsafe { enforce_abi_policy(info, access, policy)? };
let get_database = unsafe { (*access).get_database }.ok_or_else(|| {
crate::error::ExtensionError::new("get_database function pointer is null")
})?;
let db = unsafe { *get_database(info) };
let mut raw_con: duckdb_connection = core::ptr::null_mut();
let rc = unsafe { duckdb_connect(db, &raw mut raw_con) };
if rc != DuckDBSuccess {
return Err(crate::error::ExtensionError::new(
"duckdb_connect failed during extension initialization",
));
}
let con = unsafe { Connection::from_raw(raw_con, db) };
let result = catch_registration_panic(|| register(&con));
unsafe { duckdb_disconnect(&raw mut raw_con) };
result?;
Ok(true)
}
unsafe fn init_extension_internal<F>(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
api_version: &str,
policy: AbiPolicy,
register: F,
) -> Result<bool, ExtensionError>
where
F: FnOnce(duckdb_connection) -> Result<(), ExtensionError>,
{
if api_version.contains('\0') {
return Err(ExtensionError::new(
"api_version must not contain an interior NUL byte",
));
}
let have_api = unsafe {
duckdb_rs_extension_api_init(info, access, api_version)
.map_err(|e| ExtensionError::new(e.to_string()))?
};
if !have_api {
return Ok(false);
}
unsafe { enforce_abi_policy(info, access, policy)? };
let get_database = unsafe { (*access).get_database }
.ok_or_else(|| ExtensionError::new("get_database function pointer is null"))?;
let db = unsafe { *get_database(info) };
let mut raw_con: duckdb_connection = core::ptr::null_mut();
let rc = unsafe { duckdb_connect(db, &raw mut raw_con) };
if rc != DuckDBSuccess {
return Err(ExtensionError::new(
"duckdb_connect failed during extension initialization",
));
}
let result = catch_registration_panic(|| register(raw_con));
unsafe { duckdb_disconnect(&raw mut raw_con) };
result?;
Ok(true)
}
fn catch_registration_panic<F>(register: F) -> Result<(), ExtensionError>
where
F: FnOnce() -> Result<(), ExtensionError>,
{
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(register)) {
Ok(result) => result,
Err(panic) => Err(ExtensionError::new(format!(
"extension registration panicked: {}",
panic_message(&panic)
))),
}
}
fn panic_message(payload: &Box<dyn std::any::Any + Send>) -> String {
payload.downcast_ref::<&str>().map_or_else(
|| {
payload
.downcast_ref::<String>()
.map_or_else(|| String::from("<non-string panic payload>"), Clone::clone)
},
|s| (*s).to_string(),
)
}
unsafe fn enforce_abi_policy(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
policy: AbiPolicy,
) -> Result<(), ExtensionError> {
if policy == AbiPolicy::Trust {
return Ok(());
}
let check = unsafe { crate::abi::check() };
if policy == AbiPolicy::AllowUnknownEngine
&& matches!(check, crate::abi::AbiCheck::UnknownEngineVersion { .. })
{
return Ok(());
}
let Some(message) = check.error_message() else {
return Ok(());
};
if policy == AbiPolicy::Strict || policy == AbiPolicy::AllowUnknownEngine {
return Err(ExtensionError::new(message));
}
let _ = (info, access);
eprintln!("quack-rs warning: {message}");
Ok(())
}
unsafe fn report_error(
info: duckdb_extension_info,
access: *const duckdb_extension_access,
error: &ExtensionError,
) {
if access.is_null() {
return;
}
if let Some(set_error) = unsafe { (*access).set_error } {
let c_msg = error.to_c_string();
unsafe { set_error(info, c_msg.as_ptr()) };
}
}
#[cfg(test)]
mod tests {
use super::{catch_registration_panic, init_extension_internal, AbiPolicy};
use crate::error::ExtensionError;
#[test]
fn extension_error_to_c_string() {
let err = ExtensionError::new("test error message");
let cstr = err.to_c_string();
assert_eq!(cstr.to_str().unwrap(), "test error message");
}
#[test]
fn registration_success_passes_through() {
assert!(catch_registration_panic(|| Ok(())).is_ok());
}
#[test]
fn registration_error_passes_through() {
let err = catch_registration_panic(|| Err(ExtensionError::new("nope")))
.expect_err("error must propagate");
assert_eq!(err.as_str(), "nope");
}
#[test]
fn str_panic_is_converted_to_an_error() {
let err =
catch_registration_panic(|| panic!("boom")).expect_err("panic must become an error");
assert!(err.as_str().contains("registration panicked"), "{err}");
assert!(err.as_str().contains("boom"), "{err}");
}
#[test]
fn string_panic_is_converted_to_an_error() {
let err = catch_registration_panic(|| panic!("boom {}", 42))
.expect_err("panic must become an error");
assert!(err.as_str().contains("boom 42"), "{err}");
}
#[test]
fn non_string_panic_payload_still_yields_an_error() {
let err = catch_registration_panic(|| std::panic::panic_any(7u8))
.expect_err("panic must become an error");
assert!(err.as_str().contains("registration panicked"), "{err}");
}
#[test]
fn unwrap_inside_registration_does_not_escape() {
fn lookup(key: &str) -> Option<u8> {
(key == "present").then_some(1)
}
let err = catch_registration_panic(|| {
let _ = lookup("missing").unwrap();
Ok(())
})
.expect_err("unwrap panic must become an error");
assert!(err.as_str().contains("registration panicked"), "{err}");
}
#[test]
fn allow_unknown_engine_only_forgives_the_unknown_case() {
use crate::abi::AbiCheck;
let unknown = AbiCheck::UnknownEngineVersion {
engine_version: "v1.6.0".into(),
compiled_slots: 546,
};
let mismatch = AbiCheck::LayoutMismatch {
engine_version: "v1.5.0".into(),
engine_slots: 545,
compiled_slots: 546,
};
assert!(matches!(unknown, AbiCheck::UnknownEngineVersion { .. }));
assert!(!matches!(mismatch, AbiCheck::UnknownEngineVersion { .. }));
assert!(mismatch.error_message().is_some());
}
#[test]
fn api_version_with_interior_nul_is_rejected_before_ffi() {
let result = unsafe {
init_extension_internal(
core::ptr::null_mut(),
core::ptr::null(),
"v1.2\0.0",
AbiPolicy::Trust,
|_| Ok(()),
)
};
let err = result.expect_err("interior NUL must be rejected");
assert!(err.as_str().contains("NUL"), "{err}");
}
}