Skip to main content

diesel/sqlite/
auto_extension.rs

1#[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
2extern crate libsqlite3_sys as ffi;
3#[cfg(all(target_family = "wasm", target_os = "unknown"))]
4use sqlite_wasm_rs as ffi;
5
6use crate::result::Error::DatabaseError;
7use crate::result::*;
8use crate::sqlite::SqliteConnection;
9use alloc::boxed::Box;
10use alloc::string::{String, ToString};
11use core::ffi::{c_char, c_int};
12#[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
13type RawApiPointer = *const core::ffi::c_void;
14#[cfg(all(target_family = "wasm", target_os = "unknown"))]
15type RawApiPointer = *const ffi::sqlite3_api_routines;
16
17/// SQLite's auto-extension callback type.
18type RawAutoExtension = unsafe extern "C" fn(
19    db: *mut ffi::sqlite3,
20    pz_err_msg: *mut *mut c_char,
21    p_api: RawApiPointer,
22) -> c_int;
23
24// TODO: diesel 3.0 use `libsqlite3-sys` declarations after raising the minimum to 0.29.
25#[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
26mod auto_extension_ffi {
27    use super::{RawAutoExtension, c_int};
28
29    // SAFETY: These declarations match SQLite's C ABI and bypass incompatible callback types in older bindings.
30    #[allow(unsafe_code)]
31    unsafe extern "C" {
32        pub(super) fn sqlite3_auto_extension(entry_point: Option<RawAutoExtension>) -> c_int;
33        pub(super) fn sqlite3_cancel_auto_extension(entry_point: Option<RawAutoExtension>)
34        -> c_int;
35    }
36}
37
38#[cfg(all(target_family = "wasm", target_os = "unknown"))]
39use ffi as auto_extension_ffi;
40
41/// Registers an auto-extension that runs for every SQLite connection opened in
42/// this process, including non-Diesel ones.
43///
44/// This is a safe wrapper around [`sqlite3_auto_extension`][docs]. The callback
45/// receives the [`SqliteConnection`] being opened and returns `Ok(())` to
46/// continue or an error to fail the open. Use it to register SQL functions,
47/// collations, or aggregates through the usual connection API, or to initialize
48/// a statically linked C extension such as Spatialite or sqlite-vec via
49/// [`SqliteConnection::with_raw_connection`].
50///
51/// Call this before opening any connection. Extensions run in registration
52/// order, and the first error aborts the open. The callback must be a `fn` item
53/// or a closure that captures only zero-sized values (enforced at compile
54/// time), and registering the same `fn` twice is a no-op. It may run on several
55/// threads at once and must not open another connection (which would re-enter
56/// the auto-extensions and recurse) or call [`register_auto_extension`],
57/// [`cancel_auto_extension`], or [`reset_auto_extension`]. Panics are caught and
58/// turned into a failed open.
59///
60/// [docs]: https://www.sqlite.org/c3ref/auto_extension.html
61///
62/// # Example
63///
64/// ```rust
65/// use diesel::dsl::sql;
66/// use diesel::prelude::*;
67/// use diesel::sql_types::Integer;
68/// use diesel::sqlite::{register_auto_extension, reset_auto_extension, SqliteConnection};
69///
70/// // Registers a case-insensitive collation on every new connection.
71/// fn my_ext(conn: &mut SqliteConnection) -> QueryResult<()> {
72///     conn.register_collation("RUSTNOCASE", |a, b| a.to_lowercase().cmp(&b.to_lowercase()))
73/// }
74///
75/// register_auto_extension(my_ext).unwrap();
76///
77/// // Every future connection now has the collation.
78/// let mut conn = SqliteConnection::establish(":memory:").unwrap();
79/// let equal: i32 = sql::<Integer>("SELECT 'a' = 'A' COLLATE RUSTNOCASE")
80///     .get_result(&mut conn)
81///     .unwrap();
82/// assert_eq!(equal, 1);
83/// # reset_auto_extension();
84/// ```
85#[allow(unsafe_code)]
86pub fn register_auto_extension<F>(extension: F) -> QueryResult<()>
87where
88    F: Fn(&mut SqliteConnection) -> QueryResult<()> + Sync + 'static,
89{
90    // SAFETY: `entry_point` returns a stable function pointer with SQLite's callback ABI.
91    let result =
92        unsafe { auto_extension_ffi::sqlite3_auto_extension(Some(entry_point(extension))) };
93    if result == ffi::SQLITE_OK {
94        Ok(())
95    } else {
96        Err(DatabaseError(
97            DatabaseErrorKind::Unknown,
98            Box::new(ffi::code_to_str(result).to_string()),
99        ))
100    }
101}
102
103/// Removes a previously registered auto-extension, returning `true` if it was
104/// found ([docs][cancel_docs]).
105///
106/// Pass the same `fn` item given to [`register_auto_extension`]. A closure
107/// cannot be cancelled this way, because its type cannot be named again. Use
108/// [`reset_auto_extension`] to clear everything instead.
109///
110/// [cancel_docs]: https://www.sqlite.org/c3ref/cancel_auto_extension.html
111#[allow(unsafe_code)]
112pub fn cancel_auto_extension<F>(extension: F) -> bool
113where
114    F: Fn(&mut SqliteConnection) -> QueryResult<()> + Sync + 'static,
115{
116    // SAFETY: `entry_point` returns the same stable pointer used for registration.
117    unsafe { auto_extension_ffi::sqlite3_cancel_auto_extension(Some(entry_point(extension))) != 0 }
118}
119
120/// Clears **all** registered auto-extensions ([docs][reset_docs]).
121///
122/// After this call, no auto-extensions will run for newly opened connections.
123///
124/// [reset_docs]: https://www.sqlite.org/c3ref/reset_auto_extension.html
125#[allow(unsafe_code)]
126pub fn reset_auto_extension() {
127    unsafe { ffi::sqlite3_reset_auto_extension() }
128}
129
130/// Returns the trampoline for `F`. `extension` is taken by value to infer `F`,
131/// then `forget`-ten so the zero-sized callback stays conceptually alive for the
132/// process and its destructor never runs.
133fn entry_point<F>(extension: F) -> RawAutoExtension
134where
135    F: Fn(&mut SqliteConnection) -> QueryResult<()> + Sync + 'static,
136{
137    core::mem::forget(extension);
138    trampoline::<F>
139}
140
141/// The C entry point handed to SQLite, monomorphized per callback type `F` so
142/// each distinct callback maps to a distinct, stable address. SQLite's
143/// pointer-based deduplication and [`cancel_auto_extension`] rely on that.
144#[allow(unsafe_code)]
145unsafe extern "C" fn trampoline<F>(
146    db: *mut ffi::sqlite3,
147    pz_err_msg: *mut *mut c_char,
148    _p_api: RawApiPointer,
149) -> c_int
150where
151    F: Fn(&mut SqliteConnection) -> QueryResult<()> + Sync + 'static,
152{
153    const {
154        if !(core::mem::size_of::<F>() == 0) {
    {
        ::core::panicking::panic_fmt(format_args!("an auto-extension callback must not capture non-zero-sized state. Use a `fn` item or a closure that captures only zero-sized values"));
    }
};assert!(
155            core::mem::size_of::<F>() == 0,
156            "an auto-extension callback must not capture non-zero-sized state. \
157             Use a `fn` item or a closure that captures only zero-sized values"
158        );
159    }
160
161    // `_p_api` matters only for runtime-loaded shared libraries. Statically
162    // linked extensions link the SQLite symbols directly, so we ignore it.
163    let result: Result<(), String> =
164        crate::util::std_compat::catch_unwind(core::panic::AssertUnwindSafe(|| {
165            let Some(db) = core::ptr::NonNull::new(db) else {
166                return Err(String::from(
167                    "auto-extension received a null database handle",
168                ));
169            };
170            // Reconstruct a *reference* to the zero-sized callback, never an
171            // owned value, so its destructor never runs. `&F: Fn` because `F: Fn`.
172            // SAFETY: `F` is zero-sized (asserted above), so a dangling, aligned,
173            // non-null pointer is a valid `&F`, which `NonNull::dangling` provides.
174            let extension: &F = unsafe { core::ptr::NonNull::<F>::dangling().as_ref() };
175            // SAFETY: `db` is a valid handle for the duration of this call, and
176            // the borrowed connection does not take ownership of it.
177            unsafe { SqliteConnection::with_borrowed_connection(db, extension) }
178                .map_err(|e| e.to_string())
179        }))
180        .unwrap_or_else(|panic| {
181            Err(match panic_detail(panic) {
182                Some(message) => ::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!("auto-extension panicked: {0}",
                message))
    })alloc::format!("auto-extension panicked: {message}"),
183                None => String::from("auto-extension panicked"),
184            })
185        });
186
187    match result {
188        Ok(()) => ffi::SQLITE_OK,
189        Err(message) => {
190            set_error_message(pz_err_msg, &message);
191            ffi::SQLITE_ERROR
192        }
193    }
194}
195
196/// Best-effort message from a caught panic payload. The no_std `catch_unwind`
197/// carries no payload, so only the `std` variant can recover the text.
198#[cfg(feature = "std")]
199fn panic_detail(panic: alloc::boxed::Box<dyn core::any::Any + Send>) -> Option<String> {
200    panic
201        .downcast_ref::<&str>()
202        .map(|s| (*s).to_owned())
203        .or_else(|| panic.downcast_ref::<String>().cloned())
204}
205
206#[cfg(not(feature = "std"))]
207fn panic_detail(_panic: ()) -> Option<String> {
208    None
209}
210
211/// Writes `message` into `*pz_err_msg` with `sqlite3_malloc`, which is the
212/// allocator SQLite later frees it with. The message is truncated at the first
213/// NUL byte to form a valid C string, and allocation failure is ignored.
214#[allow(unsafe_code)]
215fn set_error_message(pz_err_msg: *mut *mut c_char, message: &str) {
216    if pz_err_msg.is_null() {
217        return;
218    }
219
220    let bytes = message.as_bytes();
221    let len = bytes.iter().position(|&b| b == 0).unwrap_or(bytes.len());
222
223    // SQLite sizes allocations with a C `int`. A message that does not fit is
224    // dropped rather than truncated to a bogus length.
225    let Ok(size) = c_int::try_from(len + 1) else {
226        return;
227    };
228    let buffer = unsafe { ffi::sqlite3_malloc(size) } as *mut u8;
229    if buffer.is_null() {
230        return;
231    }
232
233    unsafe {
234        core::ptr::copy_nonoverlapping(bytes.as_ptr(), buffer, len);
235        *buffer.add(len) = 0;
236        *pz_err_msg = buffer as *mut c_char;
237    }
238}
239
240// These tests either rely on sqlite calling a registered auto extension, or
241// read an error message allocated by the native library, neither of which is
242// supported when running under miri with a native libsqlite3
243// (`-Zmiri-native-lib`), so the whole module is compiled out in that case.
244#[cfg(all(test, not(miri)))]
245mod tests {
246    use super::*;
247    use crate::dsl::sql;
248    use crate::prelude::*;
249    use crate::sql_types::Integer;
250    use std::sync::Mutex;
251
252    // `sqlite3_auto_extension` is process-global, so these tests serialize on
253    // this lock and register only benign (never-failing) extensions, leaving
254    // connections opened by other tests unaffected. The failing path is covered
255    // by `trampoline_maps_result_to_return_code`, which calls the trampoline
256    // directly without touching the global registry.
257    static AUTO_EXT_TEST_LOCK: Mutex<()> = Mutex::new(());
258
259    // A benign auto-extension: registers a `TESTCOLL` collation through the
260    // normal connection API.
261    fn test_ext_init(conn: &mut SqliteConnection) -> QueryResult<()> {
262        conn.register_collation("TESTCOLL", |a, b| a.cmp(b))
263    }
264
265    fn open_memory_connection() -> SqliteConnection {
266        SqliteConnection::establish(":memory:").expect("Failed to open :memory: connection")
267    }
268
269    // Errors out if `TESTCOLL` is not registered on a freshly opened connection.
270    fn probe_collation() -> QueryResult<i32> {
271        let mut conn = open_memory_connection();
272        sql::<Integer>("SELECT 'a' = 'a' COLLATE TESTCOLL").get_result(&mut conn)
273    }
274
275    /// RAII guard that calls `reset_auto_extension()` on drop, ensuring global
276    /// state is cleaned up even if a test panics.
277    struct TestResetGuard;
278
279    impl Drop for TestResetGuard {
280        fn drop(&mut self) {
281            reset_auto_extension();
282        }
283    }
284
285    #[test]
286    fn auto_extension_lifecycle() {
287        let _lock = AUTO_EXT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
288        let _guard = TestResetGuard;
289        reset_auto_extension();
290
291        // -- 1. register + new connection has the collation --
292        register_auto_extension(test_ext_init).unwrap();
293        assert_eq!(probe_collation().unwrap(), 1);
294
295        // -- 2. cancel + new connection does NOT have the collation --
296        let removed = cancel_auto_extension(test_ext_init);
297        assert!(
298            removed,
299            "cancel should return true for registered extension"
300        );
301        assert!(
302            probe_collation().is_err(),
303            "collation should not be available after cancel"
304        );
305
306        // -- 3. cancel returns false for unregistered --
307        let removed = cancel_auto_extension(test_ext_init);
308        assert!(
309            !removed,
310            "cancel should return false for unregistered extension"
311        );
312
313        // -- 4. reset clears all --
314        register_auto_extension(test_ext_init).unwrap();
315        reset_auto_extension();
316        assert!(
317            probe_collation().is_err(),
318            "collation should not be available after reset"
319        );
320
321        // -- 5. duplicate registration is idempotent --
322        register_auto_extension(test_ext_init).unwrap();
323        register_auto_extension(test_ext_init).unwrap();
324        assert_eq!(probe_collation().unwrap(), 1);
325        // _guard drops here, ensuring reset even on panic.
326    }
327
328    // Drives the trampoline directly (via `entry_point`), without registering
329    // it in SQLite's global list, so the Ok/Err/null paths can be checked
330    // deterministically without affecting connections opened by other tests.
331    #[test]
332    #[allow(unsafe_code)]
333    fn trampoline_maps_result_to_return_code() {
334        fn ok_ext(_conn: &mut SqliteConnection) -> QueryResult<()> {
335            Ok(())
336        }
337        fn err_ext(_conn: &mut SqliteConnection) -> QueryResult<()> {
338            Err(Error::QueryBuilderError("boom".into()))
339        }
340
341        let ok_tramp = entry_point(ok_ext);
342        let err_tramp = entry_point(err_ext);
343
344        let mut conn = open_memory_connection();
345        // SAFETY: the pointer is only used for the duration of the closure,
346        // while `conn` is alive.
347        unsafe {
348            conn.with_raw_connection(|db| {
349                let mut err: *mut c_char = core::ptr::null_mut();
350
351                // Ok -> SQLITE_OK, no error message allocated.
352                let rc = ok_tramp(db, &mut err, core::ptr::null());
353                assert_eq!(rc, ffi::SQLITE_OK);
354                assert!(err.is_null());
355
356                // Err -> SQLITE_ERROR, message written via sqlite3_malloc.
357                let rc = err_tramp(db, &mut err, core::ptr::null());
358                assert_eq!(rc, ffi::SQLITE_ERROR);
359                assert!(!err.is_null());
360                let message = core::ffi::CStr::from_ptr(err)
361                    .to_string_lossy()
362                    .into_owned();
363                assert_eq!(message, "boom");
364                ffi::sqlite3_free(err as *mut core::ffi::c_void);
365
366                // Null db handle -> SQLITE_ERROR, never dereferenced.
367                let mut err: *mut c_char = core::ptr::null_mut();
368                let rc = ok_tramp(core::ptr::null_mut(), &mut err, core::ptr::null());
369                assert_eq!(rc, ffi::SQLITE_ERROR);
370                if !err.is_null() {
371                    ffi::sqlite3_free(err as *mut core::ffi::c_void);
372                }
373            })
374        }
375    }
376
377    // Regression test for the closure-reconstruction soundness fix. The callback
378    // captures a zero-sized guard with a load-bearing `Drop` that must never run,
379    // because the trampoline only reproduces the callback behind a reference and
380    // `forget`s the registered value (the old `mem::zeroed()` ran it repeatedly).
381    #[test]
382    fn callback_zero_sized_capture_is_never_dropped() {
383        use std::sync::atomic::{AtomicUsize, Ordering};
384
385        static DROPS: AtomicUsize = AtomicUsize::new(0);
386
387        // Zero-sized, `Sync`, with a load-bearing `Drop`.
388        struct Guard;
389        impl Drop for Guard {
390            fn drop(&mut self) {
391                DROPS.fetch_add(1, Ordering::SeqCst);
392            }
393        }
394
395        let _lock = AUTO_EXT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
396        let _reset = TestResetGuard;
397        reset_auto_extension();
398
399        let guard = Guard;
400        // `move` captures the zero-sized `guard` by value, so the closure is a
401        // zero-sized type with a non-trivial destructor.
402        register_auto_extension(move |conn: &mut SqliteConnection| {
403            let _ = &guard;
404            conn.register_collation("TESTCOLL", |a, b| a.cmp(b))
405        })
406        .unwrap();
407
408        for _ in 0..5 {
409            assert_eq!(probe_collation().unwrap(), 1);
410        }
411        reset_auto_extension();
412
413        assert_eq!(
414            DROPS.load(Ordering::SeqCst),
415            0,
416            "the captured guard's destructor must never run"
417        );
418    }
419
420    // A panicking callback becomes `SQLITE_ERROR` with the payload recovered
421    // into the message (`&str` and `String` payloads, generic fallback
422    // otherwise), and drives the panic-unwind path through the drop guard. Gated
423    // off WASM, where `catch_unwind` aborts because `panic = "abort"`.
424    #[test]
425    #[allow(unsafe_code)]
426    #[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
427    fn trampoline_reports_panic_message() {
428        fn panic_str(_conn: &mut SqliteConnection) -> QueryResult<()> {
429            panic!("boom-str");
430        }
431        fn panic_string(_conn: &mut SqliteConnection) -> QueryResult<()> {
432            panic!("boom-{}", 7);
433        }
434        fn panic_other(_conn: &mut SqliteConnection) -> QueryResult<()> {
435            std::panic::panic_any(7_u8);
436        }
437
438        let cases: [(RawAutoExtension, &str); 3] = [
439            (entry_point(panic_str), "auto-extension panicked: boom-str"),
440            (entry_point(panic_string), "auto-extension panicked: boom-7"),
441            (entry_point(panic_other), "auto-extension panicked"),
442        ];
443
444        let mut conn = open_memory_connection();
445        // SAFETY: the pointer is only used while `conn` is alive.
446        unsafe {
447            conn.with_raw_connection(|db| {
448                for (tramp, expected) in cases {
449                    let mut err: *mut c_char = core::ptr::null_mut();
450                    let rc = tramp(db, &mut err, core::ptr::null());
451                    assert_eq!(rc, ffi::SQLITE_ERROR);
452                    assert!(!err.is_null());
453                    let message = core::ffi::CStr::from_ptr(err)
454                        .to_string_lossy()
455                        .into_owned();
456                    assert_eq!(message, expected);
457                    ffi::sqlite3_free(err as *mut core::ffi::c_void);
458                }
459            })
460        }
461    }
462
463    #[test]
464    #[allow(unsafe_code)]
465    fn error_message_truncates_at_interior_nul() {
466        let mut err: *mut c_char = core::ptr::null_mut();
467        set_error_message(&mut err, "before\0after");
468        assert!(!err.is_null());
469        // SAFETY: `err` is a sqlite-allocated C string we own until we free it.
470        unsafe {
471            let truncated = core::ffi::CStr::from_ptr(err).to_str().unwrap();
472            assert_eq!(truncated, "before");
473            ffi::sqlite3_free(err as *mut core::ffi::c_void);
474        }
475
476        // A null out-pointer is a no-op (must not write through it or crash).
477        set_error_message(core::ptr::null_mut(), "ignored");
478    }
479}