Skip to main content

sqlitegis/sqlite/
ffi.rs

1//! SQLite extension registration via raw FFI.
2//!
3//! Registers all SQLiteGIS functions on a raw `*mut sqlite3` handle.
4//! On native targets also exports the `sqlite3_sqlitegis_init` C entry point
5//! so SQLite can load this library as a loadable extension.
6//!
7//! # Safety
8//!
9//! All `unsafe` here sits behind the SQLite `xFunc` ABI:
10//! - `argv` holds `argc` valid `*mut sqlite3_value`, and index helpers read only `< n_arg`.
11//! - blob/text borrows live only until the callback returns, so every result is copied out first.
12//! - `sqlite3_value_bytes` is read after `_blob`/`_text` so the length reflects coercion.
13//! - `extern "C"` entry points `catch_unwind`, since unwinding across the ABI is UB.
14//! - result lengths are range-checked into `c_int` via `checked_c_int_len`, buffers copied by `SQLITE_TRANSIENT`.
15
16use super::sqlite_compat::sqlite_transient;
17use super::sqlite_compat::*;
18use std::ffi::{CStr, CString};
19use std::os::raw::c_int;
20
21use crate::core::function_catalog::{
22    SqliteFunctionSpec, SQLITE_DETERMINISTIC_FUNCTIONS, SQLITE_DIRECT_ONLY_FUNCTIONS,
23};
24use crate::core::functions::accessors::*;
25use crate::core::functions::constructors::*;
26use crate::core::functions::io::*;
27use crate::core::functions::measurement::*;
28use crate::core::functions::operations::*;
29use crate::core::functions::predicates::*;
30
31// Constants
32
33const DET: c_int = SQLITE_UTF8 | SQLITE_DETERMINISTIC | SQLITE_INNOCUOUS;
34
35/// `SQLITE_DIRECTONLY` (0x80000) prevents use from triggers/views.
36/// Not yet exported by all `libsqlite3-sys` versions, so we define it here.
37const SQLITE_DIRECTONLY_FLAG: c_int = 0x0008_0000;
38const DIRECT: c_int = SQLITE_UTF8 | SQLITE_DIRECTONLY_FLAG;
39
40// Argument-extraction helpers
41
42/// Borrow argument `i` as a blob.
43///
44/// # Safety
45/// `i < argc`, and the returned slice must not outlive the callback (see the
46/// module `# Safety`).
47unsafe fn get_blob<'a>(argv: *mut *mut sqlite3_value, i: usize) -> Option<&'a [u8]> {
48    unsafe {
49        let v = *argv.add(i);
50        if sqlite3_value_type(v) == SQLITE_NULL {
51            return None;
52        }
53        // `_blob` before `_bytes` so `len` reflects coercion. from_raw_parts runs
54        // only with non-null `ptr` (zero len short-circuits above).
55        let ptr = sqlite3_value_blob(v) as *const u8;
56        let len = sqlite3_value_bytes(v) as usize;
57        if len == 0 {
58            return Some(&[]);
59        }
60        if ptr.is_null() {
61            return None;
62        }
63        Some(std::slice::from_raw_parts(ptr, len))
64    }
65}
66
67enum SqlTextArg<'a> {
68    Null,
69    Value(&'a str),
70    InvalidUtf8,
71}
72
73/// Borrow argument `i` as UTF-8 text. Same `# Safety` as [`get_blob`]. Non-UTF-8
74/// or a null pointer yields `InvalidUtf8`.
75unsafe fn get_text<'a>(argv: *mut *mut sqlite3_value, i: usize) -> SqlTextArg<'a> {
76    unsafe {
77        let v = *argv.add(i);
78        if sqlite3_value_type(v) == SQLITE_NULL {
79            return SqlTextArg::Null;
80        }
81        // `_text` before `_bytes` for a correct post-coercion `len`, null rejected
82        // before from_raw_parts.
83        let ptr = sqlite3_value_text(v);
84        let len = sqlite3_value_bytes(v) as usize;
85        if ptr.is_null() {
86            return SqlTextArg::InvalidUtf8;
87        }
88        match std::str::from_utf8(std::slice::from_raw_parts(ptr as _, len)) {
89            Ok(s) => SqlTextArg::Value(s),
90            Err(_) => SqlTextArg::InvalidUtf8,
91        }
92    }
93}
94
95enum SqlArg<T> {
96    Null,
97    Value(T),
98    InvalidType,
99}
100
101enum SqlI32Arg {
102    Null,
103    Value(i32),
104    InvalidType,
105    OutOfRange(i64),
106}
107
108unsafe fn get_f64_arg(argv: *mut *mut sqlite3_value, i: usize) -> SqlArg<f64> {
109    unsafe {
110        let v = *argv.add(i);
111        match sqlite3_value_type(v) {
112            SQLITE_NULL => SqlArg::Null,
113            SQLITE_INTEGER | SQLITE_FLOAT => SqlArg::Value(sqlite3_value_double(v)),
114            _ => SqlArg::InvalidType,
115        }
116    }
117}
118
119unsafe fn get_i32_arg(argv: *mut *mut sqlite3_value, i: usize) -> SqlI32Arg {
120    unsafe {
121        let v = *argv.add(i);
122        match sqlite3_value_type(v) {
123            SQLITE_NULL => SqlI32Arg::Null,
124            SQLITE_INTEGER => {
125                let raw = sqlite3_value_int64(v);
126                match i32::try_from(raw) {
127                    Ok(value) => SqlI32Arg::Value(value),
128                    Err(_) => SqlI32Arg::OutOfRange(raw),
129                }
130            }
131            _ => SqlI32Arg::InvalidType,
132        }
133    }
134}
135
136// Result-setting helpers
137
138fn checked_c_int_len(len: usize) -> Option<c_int> {
139    c_int::try_from(len).ok()
140}
141
142const ERROR_MSG_TOO_LARGE: &str = "internal error: error message too large";
143const PANIC_IN_CALLBACK_MSG: &str = "panic in SQLite callback";
144
145unsafe fn set_blob(ctx: *mut sqlite3_context, data: &[u8]) {
146    unsafe {
147        let Some(len) = checked_c_int_len(data.len()) else {
148            set_error(ctx, "internal error: BLOB result too large");
149            return;
150        };
151        sqlite3_result_blob(ctx, data.as_ptr().cast(), len, sqlite_transient());
152    }
153}
154
155unsafe fn set_text(ctx: *mut sqlite3_context, s: &str) {
156    unsafe {
157        let Some(len) = checked_c_int_len(s.len()) else {
158            set_error(ctx, "internal error: text result too large");
159            return;
160        };
161        sqlite3_result_text(ctx, s.as_ptr().cast(), len, sqlite_transient());
162    }
163}
164
165unsafe fn set_f64(ctx: *mut sqlite3_context, v: f64) {
166    unsafe {
167        sqlite3_result_double(ctx, v);
168    }
169}
170unsafe fn set_i64(ctx: *mut sqlite3_context, v: i64) {
171    unsafe {
172        sqlite3_result_int64(ctx, v);
173    }
174}
175unsafe fn set_i32(ctx: *mut sqlite3_context, v: i32) {
176    unsafe {
177        sqlite3_result_int(ctx, v);
178    }
179}
180unsafe fn set_null(ctx: *mut sqlite3_context) {
181    unsafe {
182        sqlite3_result_null(ctx);
183    }
184}
185
186unsafe fn set_error(ctx: *mut sqlite3_context, msg: &str) {
187    unsafe {
188        if let Some(len) = checked_c_int_len(msg.len()) {
189            sqlite3_result_error(ctx, msg.as_ptr().cast(), len);
190            return;
191        }
192
193        let len = c_int::try_from(ERROR_MSG_TOO_LARGE.len())
194            .expect("fallback error length must fit in c_int");
195        sqlite3_result_error(ctx, ERROR_MSG_TOO_LARGE.as_ptr().cast(), len);
196    }
197}
198
199unsafe fn xfunc_guard<F>(ctx: *mut sqlite3_context, label: &str, f: F)
200where
201    F: FnOnce(),
202{
203    unsafe {
204        let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(f));
205        if result.is_err() {
206            set_error(ctx, &format!("{label}: {PANIC_IN_CALLBACK_MSG}"));
207        }
208    }
209}
210
211unsafe fn require_f64_arg(
212    ctx: *mut sqlite3_context,
213    argv: *mut *mut sqlite3_value,
214    i: usize,
215    fn_name: &str,
216    arg_name: &str,
217) -> Option<f64> {
218    unsafe {
219        match get_f64_arg(argv, i) {
220            SqlArg::Value(v) => Some(v),
221            SqlArg::Null => {
222                set_null(ctx);
223                None
224            }
225            SqlArg::InvalidType => {
226                set_error(ctx, &format!("{fn_name}: {arg_name} must be numeric"));
227                None
228            }
229        }
230    }
231}
232
233unsafe fn require_i32_arg(
234    ctx: *mut sqlite3_context,
235    argv: *mut *mut sqlite3_value,
236    i: usize,
237    fn_name: &str,
238    arg_name: &str,
239) -> Option<i32> {
240    unsafe {
241        match get_i32_arg(argv, i) {
242            SqlI32Arg::Value(v) => Some(v),
243            SqlI32Arg::Null => {
244                set_null(ctx);
245                None
246            }
247            SqlI32Arg::InvalidType => {
248                set_error(ctx, &format!("{fn_name}: {arg_name} must be integer"));
249                None
250            }
251            SqlI32Arg::OutOfRange(v) => {
252                set_error(
253                    ctx,
254                    &format!("{fn_name}: {arg_name} out of range for i32: {v}"),
255                );
256                None
257            }
258        }
259    }
260}
261
262unsafe fn require_text_arg<'a>(
263    ctx: *mut sqlite3_context,
264    argv: *mut *mut sqlite3_value,
265    i: usize,
266    fn_name: &str,
267    arg_name: &str,
268) -> Option<&'a str> {
269    unsafe {
270        match get_text(argv, i) {
271            SqlTextArg::Value(v) => Some(v),
272            SqlTextArg::Null => {
273                set_null(ctx);
274                None
275            }
276            SqlTextArg::InvalidUtf8 => {
277                set_error(
278                    ctx,
279                    &format!("{fn_name}: {arg_name} must be valid UTF-8 text"),
280                );
281                None
282            }
283        }
284    }
285}
286
287unsafe fn any_arg_is_null(argv: *mut *mut sqlite3_value, arg_count: usize) -> bool {
288    unsafe {
289        for i in 0..arg_count {
290            if sqlite3_value_type(*argv.add(i)) == SQLITE_NULL {
291                return true;
292            }
293        }
294        false
295    }
296}
297
298unsafe fn optional_srid_arg(
299    ctx: *mut sqlite3_context,
300    argv: *mut *mut sqlite3_value,
301    with_srid: bool,
302    index: usize,
303    fn_name: &str,
304) -> Option<Option<i32>> {
305    unsafe {
306        if with_srid {
307            let srid = require_i32_arg(ctx, argv, index, fn_name, "srid")?;
308            Some(Some(srid))
309        } else {
310            Some(None)
311        }
312    }
313}
314
315// Convenience setter wrappers
316
317unsafe fn set_bool(ctx: *mut sqlite3_context, v: bool) {
318    unsafe {
319        set_i32(ctx, v as i32);
320    }
321}
322unsafe fn set_blob_owned(ctx: *mut sqlite3_context, v: Vec<u8>) {
323    unsafe {
324        set_blob(ctx, &v);
325    }
326}
327unsafe fn set_text_owned(ctx: *mut sqlite3_context, v: impl AsRef<str>) {
328    unsafe {
329        set_text(ctx, v.as_ref());
330    }
331}
332
333// Callback macros
334//
335// Each xfunc_* macro below generates an `unsafe extern "C" fn` with the
336// standard SQLite scalar-function signature. NULL blob/text inputs produce
337// NULL output (PostGIS-compatible). Errors produce sqlite3_result_error.
338//
339// All eleven macros share two pieces of boilerplate that are factored out
340// into the two inner helpers below:
341//
342// - `xfunc_decl!` emits the extern "C" fn signature and wraps the body in
343//   the panic-catching `xfunc_guard`.
344// - `xfunc_dispatch!` matches a `Result<T, _>` and routes the Ok arm to a
345//   setter expression while the Err arm formats and calls `set_error`.
346
347/// Emit an `unsafe extern "C" fn $name` with the SQLite scalar signature,
348/// wrapping `$body` in `xfunc_guard`. `$ctx` and `$argv` are bound to the
349/// pointer parameters so the body can name them.
350macro_rules! xfunc_decl {
351    ($name:ident, $label:expr, $ctx:ident, $argv:ident, $body:block) => {
352        unsafe extern "C" fn $name(
353            $ctx: *mut sqlite3_context,
354            _n: c_int,
355            $argv: *mut *mut sqlite3_value,
356        ) {
357            unsafe {
358                xfunc_guard($ctx, $label, || $body);
359            }
360        }
361    };
362}
363
364/// Match a `Result<T, _>` from a callback: route `Ok(v)` to `$set(ctx, v)`
365/// and format `Err(e)` into a SQLite error message tagged with `$label`.
366macro_rules! xfunc_dispatch {
367    ($ctx:expr, $label:expr, $result:expr, $set:expr) => {
368        match $result {
369            Ok(v) => $set($ctx, v),
370            Err(e) => set_error($ctx, &format!(concat!($label, ": {}"), e)),
371        }
372    };
373}
374
375/// 1 blob -> Result<T>, with a custom setter expression.
376macro_rules! xfunc_blob {
377    ($name:ident, $label:expr, $func:expr, $set:expr) => {
378        xfunc_decl!($name, $label, ctx, argv, {
379            let Some(b) = get_blob(argv, 0) else {
380                set_null(ctx);
381                return;
382            };
383            xfunc_dispatch!(ctx, $label, $func(b), $set);
384        });
385    };
386}
387
388/// 2 blobs -> Result<T>, with a custom setter expression.
389macro_rules! xfunc_blob2 {
390    ($name:ident, $label:expr, $func:expr, $set:expr) => {
391        xfunc_decl!($name, $label, ctx, argv, {
392            let Some(a) = get_blob(argv, 0) else {
393                set_null(ctx);
394                return;
395            };
396            let Some(b) = get_blob(argv, 1) else {
397                set_null(ctx);
398                return;
399            };
400            xfunc_dispatch!(ctx, $label, $func(a, b), $set);
401        });
402    };
403}
404
405/// 1 blob -> Result<Option<f64>>, where `None` maps to SQL NULL.
406///
407/// Has its own three-arm match (`Ok(Some)` / `Ok(None)` / `Err`), so it
408/// uses `xfunc_decl!` for the signature but not `xfunc_dispatch!`.
409macro_rules! xfunc_blob_opt_f64 {
410    ($name:ident, $label:expr, $func:expr) => {
411        xfunc_decl!($name, $label, ctx, argv, {
412            let Some(blob) = get_blob(argv, 0) else {
413                set_null(ctx);
414                return;
415            };
416            match $func(blob) {
417                Ok(Some(v)) => set_f64(ctx, v),
418                Ok(None) => set_null(ctx),
419                Err(e) => set_error(ctx, &format!(concat!($label, ": {}"), e)),
420            }
421        });
422    };
423}
424
425/// blob + integer arg -> Result<Vec<u8>>.
426macro_rules! xfunc_blob_i32_blob {
427    ($name:ident, $label:expr, $arg_name:expr, $func:expr) => {
428        xfunc_decl!($name, $label, ctx, argv, {
429            let Some(b) = get_blob(argv, 0) else {
430                set_null(ctx);
431                return;
432            };
433            let Some(n) = require_i32_arg(ctx, argv, 1, $label, $arg_name) else {
434                return;
435            };
436            xfunc_dispatch!(ctx, $label, ($func)(b, n), set_blob_owned);
437        });
438    };
439}
440
441/// blob + numeric arg -> Result<Vec<u8>>.
442macro_rules! xfunc_blob_f64_blob {
443    ($name:ident, $label:expr, $arg_name:expr, $func:expr) => {
444        xfunc_decl!($name, $label, ctx, argv, {
445            let Some(b) = get_blob(argv, 0) else {
446                set_null(ctx);
447                return;
448            };
449            let Some(v) = require_f64_arg(ctx, argv, 1, $label, $arg_name) else {
450                return;
451            };
452            xfunc_dispatch!(ctx, $label, ($func)(b, v), set_blob_owned);
453        });
454    };
455}
456
457/// blob + numeric arg + numeric arg -> Result<Vec<u8>>.
458macro_rules! xfunc_blob_f64_f64_blob {
459    ($name:ident, $label:expr, $arg1_name:expr, $arg2_name:expr, $func:expr) => {
460        xfunc_decl!($name, $label, ctx, argv, {
461            let Some(b) = get_blob(argv, 0) else {
462                set_null(ctx);
463                return;
464            };
465            let Some(v1) = require_f64_arg(ctx, argv, 1, $label, $arg1_name) else {
466                return;
467            };
468            let Some(v2) = require_f64_arg(ctx, argv, 2, $label, $arg2_name) else {
469                return;
470            };
471            xfunc_dispatch!(ctx, $label, ($func)(b, v1, v2), set_blob_owned);
472        });
473    };
474}
475
476/// 2 blobs + numeric arg -> Result<bool>.
477macro_rules! xfunc_blob2_f64_bool {
478    ($name:ident, $label:expr, $arg_name:expr, $func:expr) => {
479        xfunc_decl!($name, $label, ctx, argv, {
480            let Some(a) = get_blob(argv, 0) else {
481                set_null(ctx);
482                return;
483            };
484            let Some(b) = get_blob(argv, 1) else {
485                set_null(ctx);
486                return;
487            };
488            let Some(v) = require_f64_arg(ctx, argv, 2, $label, $arg_name) else {
489                return;
490            };
491            xfunc_dispatch!(ctx, $label, ($func)(a, b, v), set_bool);
492        });
493    };
494}
495
496/// 2 blobs + text arg -> Result<bool>.
497macro_rules! xfunc_blob2_text_bool {
498    ($name:ident, $label:expr, $arg_name:expr, $func:expr) => {
499        xfunc_decl!($name, $label, ctx, argv, {
500            let Some(a) = get_blob(argv, 0) else {
501                set_null(ctx);
502                return;
503            };
504            let Some(b) = get_blob(argv, 1) else {
505                set_null(ctx);
506                return;
507            };
508            let Some(v) = require_text_arg(ctx, argv, 2, $label, $arg_name) else {
509                return;
510            };
511            xfunc_dispatch!(ctx, $label, ($func)(a, b, v), set_bool);
512        });
513    };
514}
515
516/// 2 text args -> Result<bool>.
517macro_rules! xfunc_text2_bool {
518    ($name:ident, $label:expr, $arg1_name:expr, $arg2_name:expr, $func:expr) => {
519        xfunc_decl!($name, $label, ctx, argv, {
520            let Some(a) = require_text_arg(ctx, argv, 0, $label, $arg1_name) else {
521                return;
522            };
523            let Some(b) = require_text_arg(ctx, argv, 1, $label, $arg2_name) else {
524                return;
525            };
526            xfunc_dispatch!(ctx, $label, ($func)(a, b), set_bool);
527        });
528    };
529}
530
531// I/O callbacks
532
533/// text + optional SRID -> blob (generates two callbacks: 1-arg and 2-arg).
534macro_rules! xfunc_text_optsrid_blob {
535    ($name1:ident, $name2:ident, $label:expr, $func:expr) => {
536        xfunc_decl!($name1, $label, ctx, argv, {
537            let Some(t) = require_text_arg(ctx, argv, 0, $label, "wkt") else {
538                return;
539            };
540            xfunc_dispatch!(ctx, $label, $func(t, None), set_blob_owned);
541        });
542        xfunc_decl!($name2, $label, ctx, argv, {
543            let Some(t) = require_text_arg(ctx, argv, 0, $label, "wkt") else {
544                return;
545            };
546            let Some(srid) = require_i32_arg(ctx, argv, 1, $label, "srid") else {
547                return;
548            };
549            xfunc_dispatch!(ctx, $label, $func(t, Some(srid)), set_blob_owned);
550        });
551    };
552}
553
554/// blob + optional SRID -> blob (generates two callbacks: 1-arg and 2-arg).
555macro_rules! xfunc_blob_optsrid_blob {
556    ($name1:ident, $name2:ident, $label:expr, $func:expr) => {
557        xfunc_decl!($name1, $label, ctx, argv, {
558            let Some(b) = get_blob(argv, 0) else {
559                set_null(ctx);
560                return;
561            };
562            xfunc_dispatch!(ctx, $label, $func(b, None), set_blob_owned);
563        });
564        xfunc_decl!($name2, $label, ctx, argv, {
565            let Some(b) = get_blob(argv, 0) else {
566                set_null(ctx);
567                return;
568            };
569            let Some(srid) = require_i32_arg(ctx, argv, 1, $label, "srid") else {
570                return;
571            };
572            xfunc_dispatch!(ctx, $label, $func(b, Some(srid)), set_blob_owned);
573        });
574    };
575}
576
577xfunc_text_optsrid_blob!(
578    st_geomfromtext_1_xfunc,
579    st_geomfromtext_2_xfunc,
580    "ST_GeomFromText",
581    geom_from_text
582);
583xfunc_blob_optsrid_blob!(
584    st_geomfromwkb_1_xfunc,
585    st_geomfromwkb_2_xfunc,
586    "ST_GeomFromWKB",
587    geom_from_wkb
588);
589xfunc_blob!(
590    st_geomfromewkb_xfunc,
591    "ST_GeomFromEWKB",
592    geom_from_ewkb,
593    set_blob_owned
594);
595
596unsafe extern "C" fn st_geomfromgeojson_xfunc(
597    ctx: *mut sqlite3_context,
598    _n: c_int,
599    argv: *mut *mut sqlite3_value,
600) {
601    unsafe {
602        xfunc_guard(ctx, "ST_GeomFromGeoJSON", || {
603            let Some(json) = require_text_arg(ctx, argv, 0, "ST_GeomFromGeoJSON", "json") else {
604                return;
605            };
606            match geom_from_geojson(json, None) {
607                Ok(v) => set_blob(ctx, &v),
608                Err(e) => set_error(ctx, &format!("ST_GeomFromGeoJSON: {e}")),
609            }
610        });
611    }
612}
613
614xfunc_blob!(st_astext_xfunc, "ST_AsText", as_text, set_text_owned);
615xfunc_blob!(st_asewkt_xfunc, "ST_AsEWKT", as_ewkt, set_text_owned);
616xfunc_blob!(st_asbinary_xfunc, "ST_AsBinary", as_binary, set_blob_owned);
617xfunc_blob!(st_asewkb_xfunc, "ST_AsEWKB", as_ewkb, set_blob_owned);
618xfunc_blob!(
619    st_asgeojson_xfunc,
620    "ST_AsGeoJSON",
621    as_geojson,
622    set_text_owned
623);
624
625// Constructor callbacks
626
627unsafe fn st_point_impl(ctx: *mut sqlite3_context, argv: *mut *mut sqlite3_value, with_srid: bool) {
628    unsafe {
629        let arg_count = if with_srid { 3 } else { 2 };
630        if any_arg_is_null(argv, arg_count) {
631            set_null(ctx);
632            return;
633        }
634
635        let Some(x) = require_f64_arg(ctx, argv, 0, "ST_Point", "x") else {
636            return;
637        };
638        let Some(y) = require_f64_arg(ctx, argv, 1, "ST_Point", "y") else {
639            return;
640        };
641        let Some(srid) = optional_srid_arg(ctx, argv, with_srid, 2, "ST_Point") else {
642            return;
643        };
644
645        match st_point(x, y, srid) {
646            Ok(v) => set_blob(ctx, &v),
647            Err(e) => set_error(ctx, &format!("ST_Point: {e}")),
648        }
649    }
650}
651
652unsafe extern "C" fn st_point_2_xfunc(
653    ctx: *mut sqlite3_context,
654    _n: c_int,
655    argv: *mut *mut sqlite3_value,
656) {
657    unsafe {
658        xfunc_guard(ctx, "ST_Point", || {
659            st_point_impl(ctx, argv, false);
660        });
661    }
662}
663
664unsafe extern "C" fn st_point_3_xfunc(
665    ctx: *mut sqlite3_context,
666    _n: c_int,
667    argv: *mut *mut sqlite3_value,
668) {
669    unsafe {
670        xfunc_guard(ctx, "ST_Point", || {
671            st_point_impl(ctx, argv, true);
672        });
673    }
674}
675
676xfunc_blob2!(
677    st_makeline_xfunc,
678    "ST_MakeLine",
679    st_make_line,
680    set_blob_owned
681);
682xfunc_blob!(
683    st_makepolygon_xfunc,
684    "ST_MakePolygon",
685    st_make_polygon,
686    set_blob_owned
687);
688
689unsafe fn st_makeenvelope_impl(
690    ctx: *mut sqlite3_context,
691    argv: *mut *mut sqlite3_value,
692    with_srid: bool,
693) {
694    unsafe {
695        let arg_count = if with_srid { 5 } else { 4 };
696        if any_arg_is_null(argv, arg_count) {
697            set_null(ctx);
698            return;
699        }
700
701        let Some(xmin) = require_f64_arg(ctx, argv, 0, "ST_MakeEnvelope", "xmin") else {
702            return;
703        };
704        let Some(ymin) = require_f64_arg(ctx, argv, 1, "ST_MakeEnvelope", "ymin") else {
705            return;
706        };
707        let Some(xmax) = require_f64_arg(ctx, argv, 2, "ST_MakeEnvelope", "xmax") else {
708            return;
709        };
710        let Some(ymax) = require_f64_arg(ctx, argv, 3, "ST_MakeEnvelope", "ymax") else {
711            return;
712        };
713        let Some(srid) = optional_srid_arg(ctx, argv, with_srid, 4, "ST_MakeEnvelope") else {
714            return;
715        };
716
717        match st_make_envelope(xmin, ymin, xmax, ymax, srid) {
718            Ok(v) => set_blob(ctx, &v),
719            Err(e) => set_error(ctx, &format!("ST_MakeEnvelope: {e}")),
720        }
721    }
722}
723
724unsafe extern "C" fn st_makeenvelope_4_xfunc(
725    ctx: *mut sqlite3_context,
726    _n: c_int,
727    argv: *mut *mut sqlite3_value,
728) {
729    unsafe {
730        xfunc_guard(ctx, "ST_MakeEnvelope", || {
731            st_makeenvelope_impl(ctx, argv, false);
732        });
733    }
734}
735
736unsafe extern "C" fn st_makeenvelope_5_xfunc(
737    ctx: *mut sqlite3_context,
738    _n: c_int,
739    argv: *mut *mut sqlite3_value,
740) {
741    unsafe {
742        xfunc_guard(ctx, "ST_MakeEnvelope", || {
743            st_makeenvelope_impl(ctx, argv, true);
744        });
745    }
746}
747
748xfunc_blob2!(st_collect_xfunc, "ST_Collect", st_collect, set_blob_owned);
749
750unsafe extern "C" fn st_tileenvelope_xfunc(
751    ctx: *mut sqlite3_context,
752    _n: c_int,
753    argv: *mut *mut sqlite3_value,
754) {
755    unsafe {
756        xfunc_guard(ctx, "ST_TileEnvelope", || {
757            let Some(zoom_i32) = require_i32_arg(ctx, argv, 0, "ST_TileEnvelope", "zoom") else {
758                return;
759            };
760            let Some(tile_x_i32) = require_i32_arg(ctx, argv, 1, "ST_TileEnvelope", "tile x")
761            else {
762                return;
763            };
764            let Some(tile_y_i32) = require_i32_arg(ctx, argv, 2, "ST_TileEnvelope", "tile y")
765            else {
766                return;
767            };
768
769            if zoom_i32 < 0 {
770                set_error(ctx, "ST_TileEnvelope: zoom must be non-negative");
771                return;
772            }
773            if tile_x_i32 < 0 {
774                set_error(ctx, "ST_TileEnvelope: tile x must be non-negative");
775                return;
776            }
777            if tile_y_i32 < 0 {
778                set_error(ctx, "ST_TileEnvelope: tile y must be non-negative");
779                return;
780            }
781
782            let zoom = zoom_i32 as u32;
783            let tile_x = tile_x_i32 as u32;
784            let tile_y = tile_y_i32 as u32;
785            match st_tile_envelope(zoom, tile_x, tile_y) {
786                Ok(v) => set_blob(ctx, &v),
787                Err(e) => set_error(ctx, &format!("ST_TileEnvelope: {e}")),
788            }
789        });
790    }
791}
792
793// Accessor callbacks
794
795xfunc_blob!(st_srid_xfunc, "ST_SRID", st_srid, set_i32);
796
797unsafe extern "C" fn st_setsrid_xfunc(
798    ctx: *mut sqlite3_context,
799    _n: c_int,
800    argv: *mut *mut sqlite3_value,
801) {
802    unsafe {
803        xfunc_guard(ctx, "ST_SetSRID", || {
804            let Some(b) = get_blob(argv, 0) else {
805                set_null(ctx);
806                return;
807            };
808            let Some(srid) = require_i32_arg(ctx, argv, 1, "ST_SetSRID", "srid") else {
809                return;
810            };
811            match st_set_srid(b, srid) {
812                Ok(v) => set_blob(ctx, &v),
813                Err(e) => set_error(ctx, &format!("ST_SetSRID: {e}")),
814            }
815        });
816    }
817}
818
819xfunc_blob!(
820    st_geometrytype_xfunc,
821    "ST_GeometryType",
822    st_geometry_type,
823    set_text_owned
824);
825xfunc_blob!(st_ndims_xfunc, "ST_NDims", st_ndims, set_i32);
826xfunc_blob!(st_coorddim_xfunc, "ST_CoordDim", st_coord_dim, set_i32);
827xfunc_blob!(st_zmflag_xfunc, "ST_Zmflag", st_zmflag, set_i32);
828xfunc_blob!(st_isempty_xfunc, "ST_IsEmpty", st_is_empty, set_bool);
829xfunc_blob!(st_memsize_xfunc, "ST_MemSize", st_mem_size, set_i64);
830xfunc_blob_opt_f64!(st_x_xfunc, "ST_X", st_x);
831xfunc_blob_opt_f64!(st_y_xfunc, "ST_Y", st_y);
832xfunc_blob_opt_f64!(st_z_xfunc, "ST_Z", st_z);
833
834xfunc_blob!(st_numpoints_xfunc, "ST_NumPoints", st_num_points, set_i32);
835xfunc_blob!(st_npoints_xfunc, "ST_NPoints", st_npoints, set_i32);
836xfunc_blob!(
837    st_numgeometries_xfunc,
838    "ST_NumGeometries",
839    st_num_geometries,
840    set_i32
841);
842xfunc_blob!(
843    st_numinteriorrings_xfunc,
844    "ST_NumInteriorRings",
845    st_num_interior_rings,
846    set_i32
847);
848xfunc_blob!(st_numrings_xfunc, "ST_NumRings", st_num_rings, set_i32);
849xfunc_blob_i32_blob!(st_pointn_xfunc, "ST_PointN", "n", |b, n| st_point_n(
850    b, n, None
851));
852
853xfunc_blob!(
854    st_startpoint_xfunc,
855    "ST_StartPoint",
856    st_start_point,
857    set_blob_owned
858);
859xfunc_blob!(
860    st_endpoint_xfunc,
861    "ST_EndPoint",
862    st_end_point,
863    set_blob_owned
864);
865xfunc_blob!(
866    st_exteriorring_xfunc,
867    "ST_ExteriorRing",
868    st_exterior_ring,
869    set_blob_owned
870);
871xfunc_blob_i32_blob!(
872    st_interiorringn_xfunc,
873    "ST_InteriorRingN",
874    "n",
875    st_interior_ring_n
876);
877xfunc_blob_i32_blob!(st_geometryn_xfunc, "ST_GeometryN", "n", st_geometry_n);
878
879xfunc_blob!(st_dimension_xfunc, "ST_Dimension", st_dimension, set_i32);
880xfunc_blob!(
881    st_envelope_xfunc,
882    "ST_Envelope",
883    st_envelope,
884    set_blob_owned
885);
886xfunc_blob!(st_isvalid_xfunc, "ST_IsValid", st_is_valid, set_bool);
887xfunc_blob!(
888    st_isvalidreason_xfunc,
889    "ST_IsValidReason",
890    st_is_valid_reason,
891    set_text_owned
892);
893
894// Measurement callbacks
895
896xfunc_blob!(st_area_xfunc, "ST_Area", st_area, set_f64);
897xfunc_blob!(st_length_xfunc, "ST_Length", st_length, set_f64);
898xfunc_blob!(st_perimeter_xfunc, "ST_Perimeter", st_perimeter, set_f64);
899xfunc_blob2!(st_distance_xfunc, "ST_Distance", st_distance, set_f64);
900xfunc_blob!(
901    st_centroid_xfunc,
902    "ST_Centroid",
903    st_centroid,
904    set_blob_owned
905);
906xfunc_blob!(
907    st_pointonsurface_xfunc,
908    "ST_PointOnSurface",
909    st_point_on_surface,
910    set_blob_owned
911);
912xfunc_blob2!(
913    st_hausdorffdistance_xfunc,
914    "ST_HausdorffDistance",
915    st_hausdorff_distance,
916    set_f64
917);
918xfunc_blob_opt_f64!(st_xmin_xfunc, "ST_XMin", st_xmin);
919xfunc_blob_opt_f64!(st_xmax_xfunc, "ST_XMax", st_xmax);
920xfunc_blob_opt_f64!(st_ymin_xfunc, "ST_YMin", st_ymin);
921xfunc_blob_opt_f64!(st_ymax_xfunc, "ST_YMax", st_ymax);
922xfunc_blob2!(
923    st_distancesphere_xfunc,
924    "ST_DistanceSphere",
925    st_distance_sphere,
926    set_f64
927);
928xfunc_blob2!(
929    st_distancespheroid_xfunc,
930    "ST_DistanceSpheroid",
931    st_distance_spheroid,
932    set_f64
933);
934xfunc_blob!(
935    st_lengthsphere_xfunc,
936    "ST_LengthSphere",
937    st_length_sphere,
938    set_f64
939);
940xfunc_blob2!(st_azimuth_xfunc, "ST_Azimuth", st_azimuth, set_f64);
941
942xfunc_blob_f64_f64_blob!(
943    st_project_xfunc,
944    "ST_Project",
945    "distance",
946    "azimuth",
947    st_project
948);
949
950xfunc_blob2!(
951    st_closestpoint_xfunc,
952    "ST_ClosestPoint",
953    st_closest_point,
954    set_blob_owned
955);
956
957// Operation callbacks
958
959xfunc_blob2!(st_union_xfunc, "ST_Union", st_union, set_blob_owned);
960xfunc_blob2!(
961    st_intersection_xfunc,
962    "ST_Intersection",
963    st_intersection,
964    set_blob_owned
965);
966xfunc_blob2!(
967    st_difference_xfunc,
968    "ST_Difference",
969    st_difference,
970    set_blob_owned
971);
972xfunc_blob2!(
973    st_symdifference_xfunc,
974    "ST_SymDifference",
975    st_sym_difference,
976    set_blob_owned
977);
978
979xfunc_blob_f64_blob!(st_buffer_xfunc, "ST_Buffer", "distance", st_buffer);
980
981// Predicate callbacks
982
983xfunc_blob2!(
984    st_intersects_xfunc,
985    "ST_Intersects",
986    st_intersects,
987    set_bool
988);
989xfunc_blob2!(st_contains_xfunc, "ST_Contains", st_contains, set_bool);
990xfunc_blob2!(st_within_xfunc, "ST_Within", st_within, set_bool);
991xfunc_blob2!(st_disjoint_xfunc, "ST_Disjoint", st_disjoint, set_bool);
992
993xfunc_blob2_f64_bool!(st_dwithin_xfunc, "ST_DWithin", "distance", st_dwithin);
994xfunc_blob2_f64_bool!(
995    st_dwithinsphere_xfunc,
996    "ST_DWithinSphere",
997    "distance",
998    st_dwithin_sphere
999);
1000xfunc_blob2_f64_bool!(
1001    st_dwithinspheroid_xfunc,
1002    "ST_DWithinSpheroid",
1003    "distance",
1004    st_dwithin_spheroid
1005);
1006
1007xfunc_blob2!(st_covers_xfunc, "ST_Covers", st_covers, set_bool);
1008xfunc_blob2!(st_coveredby_xfunc, "ST_CoveredBy", st_covered_by, set_bool);
1009xfunc_blob2!(st_equals_xfunc, "ST_Equals", st_equals, set_bool);
1010xfunc_blob2!(st_touches_xfunc, "ST_Touches", st_touches, set_bool);
1011xfunc_blob2!(st_crosses_xfunc, "ST_Crosses", st_crosses, set_bool);
1012xfunc_blob2!(st_overlaps_xfunc, "ST_Overlaps", st_overlaps, set_bool);
1013
1014xfunc_blob2!(st_relate_2_xfunc, "ST_Relate", st_relate, set_text_owned);
1015
1016xfunc_blob2_text_bool!(
1017    st_relate_3_xfunc,
1018    "ST_Relate",
1019    "pattern",
1020    st_relate_match_geoms
1021);
1022xfunc_text2_bool!(
1023    st_relatematch_xfunc,
1024    "ST_RelateMatch",
1025    "matrix",
1026    "pattern",
1027    st_relate_match
1028);
1029
1030// Spatial index helpers
1031
1032fn validate_identifier(s: &str) -> Option<&str> {
1033    if s.is_empty() {
1034        return None;
1035    }
1036    if s.bytes().all(|b| b.is_ascii_alphanumeric() || b == b'_') {
1037        Some(s)
1038    } else {
1039        None
1040    }
1041}
1042
1043fn sql_to_cstring(sql: &str) -> std::result::Result<CString, std::ffi::NulError> {
1044    CString::new(sql)
1045}
1046
1047unsafe fn exec_sql_inner(db: *mut sqlite3, sql: &str, ctx: Option<*mut sqlite3_context>) -> c_int {
1048    unsafe {
1049        let c_sql = match sql_to_cstring(sql) {
1050            Ok(v) => v,
1051            Err(_) => {
1052                if let Some(ctx) = ctx {
1053                    set_error(ctx, "internal error: generated SQL contains NUL byte");
1054                }
1055                return SQLITE_ERROR;
1056            }
1057        };
1058
1059        let mut err_msg: *mut std::ffi::c_char = std::ptr::null_mut();
1060        let rc = sqlite3_exec(db, c_sql.as_ptr(), None, std::ptr::null_mut(), &mut err_msg);
1061
1062        if rc != SQLITE_OK {
1063            if let Some(ctx) = ctx {
1064                if err_msg.is_null() {
1065                    set_error(ctx, "exec_sql failed");
1066                } else {
1067                    let msg = CStr::from_ptr(err_msg).to_string_lossy();
1068                    set_error(ctx, &msg);
1069                }
1070            }
1071        }
1072
1073        if !err_msg.is_null() {
1074            sqlite3_free(err_msg.cast());
1075        }
1076        rc
1077    }
1078}
1079
1080/// Run SQL via `sqlite3_exec`, returning `SQLITE_OK` on success.
1081/// On failure, sets `sqlite3_result_error` on `ctx` with the error message
1082/// from SQLite and frees it via `sqlite3_free`.
1083unsafe fn exec_sql(db: *mut sqlite3, ctx: *mut sqlite3_context, sql: &str) -> c_int {
1084    unsafe { exec_sql_inner(db, sql, Some(ctx)) }
1085}
1086
1087/// Run SQL via `sqlite3_exec` but never touch sqlite3_result_error.
1088/// Used for best-effort rollback paths where the original error should win.
1089unsafe fn exec_sql_silent(db: *mut sqlite3, sql: &str) -> c_int {
1090    unsafe { exec_sql_inner(db, sql, None) }
1091}
1092
1093unsafe fn rollback_savepoint(db: *mut sqlite3, ctx: *mut sqlite3_context, savepoint: &str) {
1094    unsafe {
1095        let _ = ctx;
1096        let _ = exec_sql_silent(db, &format!("ROLLBACK TO {savepoint}"));
1097        let _ = exec_sql_silent(db, &format!("RELEASE {savepoint}"));
1098    }
1099}
1100
1101unsafe fn sqlite_master_lookup_text(
1102    db: *mut sqlite3,
1103    sql: &str,
1104) -> std::result::Result<Option<String>, String> {
1105    unsafe {
1106        let c_sql = sql_to_cstring(sql)
1107            .map_err(|_| "internal error: generated SQL contains NUL byte".to_string())?;
1108        let mut stmt: *mut sqlite3_stmt = std::ptr::null_mut();
1109        let rc = sqlite3_prepare_v2(db, c_sql.as_ptr(), -1, &mut stmt, std::ptr::null_mut());
1110        if rc != SQLITE_OK {
1111            return Err(CStr::from_ptr(sqlite3_errmsg(db))
1112                .to_string_lossy()
1113                .into_owned());
1114        }
1115
1116        let mut result = None;
1117        let step = sqlite3_step(stmt);
1118        if step == SQLITE_ROW {
1119            if sqlite3_column_type(stmt, 0) != SQLITE_NULL {
1120                let ptr = sqlite3_column_text(stmt, 0);
1121                if !ptr.is_null() {
1122                    result = Some(CStr::from_ptr(ptr.cast()).to_string_lossy().into_owned());
1123                }
1124            }
1125        } else if step != SQLITE_DONE {
1126            let err = CStr::from_ptr(sqlite3_errmsg(db))
1127                .to_string_lossy()
1128                .into_owned();
1129            let _ = sqlite3_finalize(stmt);
1130            return Err(err);
1131        }
1132
1133        let _ = sqlite3_finalize(stmt);
1134        Ok(result)
1135    }
1136}
1137
1138const SPATIAL_INDEX_CATALOG_TABLE: &str = "sqlitegis_spatial_index_catalog";
1139const SPATIAL_INDEX_CATALOG_REQUIRED_COLUMNS: [&str; 3] = ["prefix", "table_name", "column_name"];
1140
1141#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1142enum SpatialIndexOwnership {
1143    Owned,
1144    Absent,
1145}
1146
1147unsafe fn lookup_sqlite_master_object_type(
1148    db: *mut sqlite3,
1149    object_name: &str,
1150) -> std::result::Result<Option<String>, String> {
1151    unsafe {
1152        let sql = format!("SELECT type FROM sqlite_master WHERE name = '{object_name}' LIMIT 1");
1153        sqlite_master_lookup_text(db, &sql)
1154    }
1155}
1156
1157unsafe fn inspect_spatial_index_catalog_columns(
1158    db: *mut sqlite3,
1159) -> std::result::Result<(bool, bool, bool), String> {
1160    unsafe {
1161        let sql = format!("PRAGMA table_info([{SPATIAL_INDEX_CATALOG_TABLE}])");
1162        let c_sql = sql_to_cstring(&sql)
1163            .map_err(|_| "internal error: generated SQL contains NUL byte".to_string())?;
1164        let mut stmt: *mut sqlite3_stmt = std::ptr::null_mut();
1165        let rc = sqlite3_prepare_v2(db, c_sql.as_ptr(), -1, &mut stmt, std::ptr::null_mut());
1166        if rc != SQLITE_OK {
1167            return Err(CStr::from_ptr(sqlite3_errmsg(db))
1168                .to_string_lossy()
1169                .into_owned());
1170        }
1171
1172        let mut has_prefix = false;
1173        let mut has_table_name = false;
1174        let mut has_column_name = false;
1175
1176        loop {
1177            let step = sqlite3_step(stmt);
1178            if step == SQLITE_ROW {
1179                if sqlite3_column_type(stmt, 1) != SQLITE_NULL {
1180                    let ptr = sqlite3_column_text(stmt, 1);
1181                    if !ptr.is_null() {
1182                        let column_name = CStr::from_ptr(ptr.cast()).to_string_lossy();
1183                        match column_name.as_ref() {
1184                            "prefix" => has_prefix = true,
1185                            "table_name" => has_table_name = true,
1186                            "column_name" => has_column_name = true,
1187                            _ => {}
1188                        }
1189                    }
1190                }
1191                continue;
1192            }
1193            if step == SQLITE_DONE {
1194                break;
1195            }
1196            let err = CStr::from_ptr(sqlite3_errmsg(db))
1197                .to_string_lossy()
1198                .into_owned();
1199            let _ = sqlite3_finalize(stmt);
1200            return Err(err);
1201        }
1202
1203        let _ = sqlite3_finalize(stmt);
1204        Ok((has_prefix, has_table_name, has_column_name))
1205    }
1206}
1207
1208unsafe fn validate_spatial_index_catalog_shape(
1209    db: *mut sqlite3,
1210    ctx: *mut sqlite3_context,
1211    label: &str,
1212) -> bool {
1213    unsafe {
1214        let object_type = match lookup_sqlite_master_object_type(db, SPATIAL_INDEX_CATALOG_TABLE) {
1215            Ok(v) => v,
1216            Err(e) => {
1217                set_error(
1218                    ctx,
1219                    &format!("{label}: failed to inspect spatial index catalog metadata: {e}"),
1220                );
1221                return false;
1222            }
1223        };
1224        let Some(object_type) = object_type else {
1225            set_error(
1226                ctx,
1227                &format!(
1228                    "{label}: failed to inspect spatial index catalog metadata: \
1229                 missing sqlite_master entry for [{SPATIAL_INDEX_CATALOG_TABLE}]"
1230                ),
1231            );
1232            return false;
1233        };
1234        if object_type != "table" {
1235            set_error(
1236                ctx,
1237                &format!(
1238                    "{label}: invalid spatial index catalog object type for \
1239                 [{SPATIAL_INDEX_CATALOG_TABLE}] (expected table, found [{object_type}])"
1240                ),
1241            );
1242            return false;
1243        }
1244
1245        let (has_prefix, has_table_name, has_column_name) =
1246            match inspect_spatial_index_catalog_columns(db) {
1247                Ok(v) => v,
1248                Err(e) => {
1249                    set_error(
1250                        ctx,
1251                        &format!("{label}: failed to inspect spatial index catalog metadata: {e}"),
1252                    );
1253                    return false;
1254                }
1255            };
1256
1257        let present = [has_prefix, has_table_name, has_column_name];
1258        for (i, required_column) in SPATIAL_INDEX_CATALOG_REQUIRED_COLUMNS.iter().enumerate() {
1259            if !present[i] {
1260                set_error(
1261                    ctx,
1262                    &format!(
1263                        "{label}: invalid spatial index catalog schema for \
1264                     [{SPATIAL_INDEX_CATALOG_TABLE}] (missing required column [{required_column}])"
1265                    ),
1266                );
1267                return false;
1268            }
1269        }
1270
1271        true
1272    }
1273}
1274
1275unsafe fn ensure_spatial_index_catalog_table(
1276    db: *mut sqlite3,
1277    ctx: *mut sqlite3_context,
1278    label: &str,
1279) -> bool {
1280    unsafe {
1281        let object_type = match lookup_sqlite_master_object_type(db, SPATIAL_INDEX_CATALOG_TABLE) {
1282            Ok(v) => v,
1283            Err(e) => {
1284                set_error(
1285                    ctx,
1286                    &format!("{label}: failed to inspect spatial index catalog metadata: {e}"),
1287                );
1288                return false;
1289            }
1290        };
1291        if let Some(object_type) = object_type {
1292            if object_type != "table" {
1293                set_error(
1294                    ctx,
1295                    &format!(
1296                        "{label}: invalid spatial index catalog object type for \
1297                     [{SPATIAL_INDEX_CATALOG_TABLE}] (expected table, found [{object_type}])"
1298                    ),
1299                );
1300                return false;
1301            }
1302        }
1303
1304        let sql = format!(
1305            "CREATE TABLE IF NOT EXISTS [{SPATIAL_INDEX_CATALOG_TABLE}] (\
1306         prefix TEXT PRIMARY KEY, \
1307         table_name TEXT NOT NULL, \
1308         column_name TEXT NOT NULL, \
1309         UNIQUE(table_name, column_name)\
1310         )"
1311        );
1312        if exec_sql_silent(db, &sql) == SQLITE_OK {
1313            return true;
1314        }
1315
1316        let err = CStr::from_ptr(sqlite3_errmsg(db))
1317            .to_string_lossy()
1318            .into_owned();
1319        set_error(
1320            ctx,
1321            &format!("{label}: failed to ensure spatial index catalog: {err}"),
1322        );
1323        false
1324    }
1325}
1326
1327unsafe fn lookup_spatial_index_catalog_owner(
1328    db: *mut sqlite3,
1329    prefix: &str,
1330) -> std::result::Result<Option<(String, String)>, String> {
1331    unsafe {
1332        let sql = format!(
1333            "SELECT table_name FROM [{SPATIAL_INDEX_CATALOG_TABLE}] \
1334         WHERE prefix = '{prefix}' LIMIT 1"
1335        );
1336        let owner_table = sqlite_master_lookup_text(db, &sql)?;
1337        let Some(owner_table) = owner_table else {
1338            return Ok(None);
1339        };
1340
1341        let sql = format!(
1342            "SELECT column_name FROM [{SPATIAL_INDEX_CATALOG_TABLE}] \
1343         WHERE prefix = '{prefix}' LIMIT 1"
1344        );
1345        let owner_column = sqlite_master_lookup_text(db, &sql)?;
1346        let Some(owner_column) = owner_column else {
1347            return Err(format!(
1348                "internal error: catalog row for prefix [{prefix}] is missing column_name"
1349            ));
1350        };
1351        Ok(Some((owner_table, owner_column)))
1352    }
1353}
1354
1355unsafe fn managed_spatial_index_objects_exist(
1356    db: *mut sqlite3,
1357    prefix: &str,
1358) -> std::result::Result<bool, String> {
1359    unsafe {
1360        let rtree_name = format!("{prefix}_rtree");
1361        let sql = format!(
1362            "SELECT name FROM sqlite_master WHERE name IN (\
1363         '{rtree_name}', \
1364         '{rtree_name}_node', \
1365         '{rtree_name}_parent', \
1366         '{rtree_name}_rowid', \
1367         '{prefix}_insert', \
1368         '{prefix}_update', \
1369         '{prefix}_delete'\
1370         ) LIMIT 1"
1371        );
1372        Ok(sqlite_master_lookup_text(db, &sql)?.is_some())
1373    }
1374}
1375
1376unsafe fn ensure_spatial_index_table_shape(
1377    db: *mut sqlite3,
1378    ctx: *mut sqlite3_context,
1379    prefix: &str,
1380    label: &str,
1381) -> bool {
1382    unsafe {
1383        // The managed object `{prefix}_rtree` must be either absent or a real
1384        // SQLite table backed by the expected R-tree shadow tables.
1385        let rtree_name = format!("{prefix}_rtree");
1386        let sql = format!("SELECT type FROM sqlite_master WHERE name = '{rtree_name}' LIMIT 1");
1387        let object_type = match sqlite_master_lookup_text(db, &sql) {
1388            Ok(v) => v,
1389            Err(e) => {
1390                set_error(
1391                    ctx,
1392                    &format!("{label}: failed to inspect sqlite_master: {e}"),
1393                );
1394                return false;
1395            }
1396        };
1397        if let Some(object_type) = object_type {
1398            if object_type != "table" {
1399                set_error(
1400                    ctx,
1401                    &format!(
1402                        "{label}: unexpected sqlite_master entry for [{rtree_name}] \
1403                     (type [{object_type}]); expected table"
1404                    ),
1405                );
1406                return false;
1407            }
1408
1409            for shadow_suffix in &["_node", "_parent", "_rowid"] {
1410                let shadow_name = format!("{rtree_name}{shadow_suffix}");
1411                let sql = format!(
1412                    "SELECT name FROM sqlite_master \
1413                 WHERE type = 'table' AND name = '{shadow_name}' LIMIT 1"
1414                );
1415                let shadow_exists = match sqlite_master_lookup_text(db, &sql) {
1416                    Ok(v) => v,
1417                    Err(e) => {
1418                        set_error(
1419                            ctx,
1420                            &format!("{label}: failed to inspect sqlite_master: {e}"),
1421                        );
1422                        return false;
1423                    }
1424                };
1425                if shadow_exists.is_none() {
1426                    set_error(
1427                        ctx,
1428                        &format!(
1429                            "{label}: existing table [{rtree_name}] is not an R-tree index \
1430                         managed by SQLiteGIS (missing shadow table [{shadow_name}])"
1431                        ),
1432                    );
1433                    return false;
1434                }
1435            }
1436        }
1437
1438        true
1439    }
1440}
1441
1442unsafe fn ensure_spatial_index_objects_owned_by_table(
1443    db: *mut sqlite3,
1444    ctx: *mut sqlite3_context,
1445    table: &str,
1446    column: &str,
1447    label: &str,
1448) -> Option<SpatialIndexOwnership> {
1449    unsafe {
1450        // Object names are built as `{table}_{column}_...`. Different input pairs
1451        // can collide (e.g. `a_b`+`c` vs `a`+`b_c`). Detect and fail fast instead
1452        // of silently reusing another table's triggers/index objects.
1453        let prefix = format!("{table}_{column}");
1454        for suffix in &["_insert", "_update", "_delete"] {
1455            let trigger_name = format!("{prefix}{suffix}");
1456            let sql = format!(
1457                "SELECT tbl_name FROM sqlite_master \
1458             WHERE type = 'trigger' AND name = '{trigger_name}' LIMIT 1"
1459            );
1460            let owner = match sqlite_master_lookup_text(db, &sql) {
1461                Ok(v) => v,
1462                Err(e) => {
1463                    set_error(
1464                        ctx,
1465                        &format!("{label}: failed to inspect sqlite_master: {e}"),
1466                    );
1467                    return None;
1468                }
1469            };
1470            if let Some(owner) = owner {
1471                if owner != table {
1472                    set_error(
1473                        ctx,
1474                        &format!(
1475                            "{label}: naming collision for trigger [{trigger_name}] \
1476                         between tables [{owner}] and [{table}]"
1477                        ),
1478                    );
1479                    return None;
1480                }
1481            }
1482        }
1483
1484        if !ensure_spatial_index_table_shape(db, ctx, &prefix, label) {
1485            return None;
1486        }
1487
1488        let owner = match lookup_spatial_index_catalog_owner(db, &prefix) {
1489            Ok(v) => v,
1490            Err(e) => {
1491                set_error(ctx, &format!("{label}: failed to inspect catalog: {e}"));
1492                return None;
1493            }
1494        };
1495        if let Some((owner_table, owner_column)) = owner {
1496            if owner_table == table && owner_column == column {
1497                return Some(SpatialIndexOwnership::Owned);
1498            }
1499            set_error(
1500                ctx,
1501                &format!(
1502                    "{label}: naming collision for managed prefix [{prefix}] \
1503                 between [{owner_table}.{owner_column}] and [{table}.{column}]"
1504                ),
1505            );
1506            return None;
1507        }
1508
1509        let objects_exist = match managed_spatial_index_objects_exist(db, &prefix) {
1510            Ok(v) => v,
1511            Err(e) => {
1512                set_error(
1513                    ctx,
1514                    &format!("{label}: failed to inspect sqlite_master: {e}"),
1515                );
1516                return None;
1517            }
1518        };
1519        if objects_exist {
1520            set_error(
1521                ctx,
1522                &format!(
1523                    "{label}: cannot prove ownership for [{prefix}] because managed objects exist \
1524                 without an ownership marker"
1525                ),
1526            );
1527            return None;
1528        }
1529
1530        Some(SpatialIndexOwnership::Absent)
1531    }
1532}
1533
1534// Spatial index callbacks
1535
1536/// Extract and validate `(table, column)` identifiers from the first two args.
1537/// On failure, sets an error on `ctx` and returns `None`.
1538unsafe fn get_table_column<'a>(
1539    ctx: *mut sqlite3_context,
1540    argv: *mut *mut sqlite3_value,
1541    label: &str,
1542) -> Option<(&'a str, &'a str)> {
1543    unsafe {
1544        let table = match get_text(argv, 0) {
1545            SqlTextArg::Value(v) => v,
1546            SqlTextArg::Null => {
1547                set_error(ctx, &format!("{label}: table name must not be NULL"));
1548                return None;
1549            }
1550            SqlTextArg::InvalidUtf8 => {
1551                set_error(
1552                    ctx,
1553                    &format!("{label}: table name must be valid UTF-8 text"),
1554                );
1555                return None;
1556            }
1557        };
1558        let column = match get_text(argv, 1) {
1559            SqlTextArg::Value(v) => v,
1560            SqlTextArg::Null => {
1561                set_error(ctx, &format!("{label}: column name must not be NULL"));
1562                return None;
1563            }
1564            SqlTextArg::InvalidUtf8 => {
1565                set_error(
1566                    ctx,
1567                    &format!("{label}: column name must be valid UTF-8 text"),
1568                );
1569                return None;
1570            }
1571        };
1572        let Some(table) = validate_identifier(table) else {
1573            set_error(
1574                ctx,
1575                &format!("{label}: invalid table name (only [a-zA-Z0-9_] allowed)"),
1576            );
1577            return None;
1578        };
1579        let Some(column) = validate_identifier(column) else {
1580            set_error(
1581                ctx,
1582                &format!("{label}: invalid column name (only [a-zA-Z0-9_] allowed)"),
1583            );
1584            return None;
1585        };
1586        Some((table, column))
1587    }
1588}
1589
1590unsafe extern "C" fn create_spatial_index_xfunc(
1591    ctx: *mut sqlite3_context,
1592    _n: c_int,
1593    argv: *mut *mut sqlite3_value,
1594) {
1595    unsafe {
1596        xfunc_guard(ctx, "CreateSpatialIndex", || {
1597            let Some((table, column)) = get_table_column(ctx, argv, "CreateSpatialIndex") else {
1598                return;
1599            };
1600
1601            let db = sqlite3_context_db_handle(ctx);
1602            let prefix = format!("{table}_{column}");
1603            let rtree = format!("{prefix}_rtree");
1604            let savepoint = "sqlitegis_create_spatial_index";
1605
1606            if exec_sql(db, ctx, &format!("SAVEPOINT {savepoint}")) != SQLITE_OK {
1607                return;
1608            }
1609
1610            if !ensure_spatial_index_catalog_table(db, ctx, "CreateSpatialIndex") {
1611                rollback_savepoint(db, ctx, savepoint);
1612                return;
1613            }
1614            if !validate_spatial_index_catalog_shape(db, ctx, "CreateSpatialIndex") {
1615                rollback_savepoint(db, ctx, savepoint);
1616                return;
1617            }
1618
1619            if ensure_spatial_index_objects_owned_by_table(
1620                db,
1621                ctx,
1622                table,
1623                column,
1624                "CreateSpatialIndex",
1625            )
1626            .is_none()
1627            {
1628                rollback_savepoint(db, ctx, savepoint);
1629                return;
1630            }
1631
1632            // Reject WITHOUT ROWID tables before creating any state. The
1633            // maintenance triggers reference NEW.rowid and OLD.rowid, which only
1634            // exist on regular rowid tables. SQLite refuses to prepare a SELECT
1635            // of rowid against a WITHOUT ROWID table, so the probe fails cleanly
1636            // at parse time. A successful probe also incidentally proves the
1637            // table exists.
1638            let probe = format!("SELECT rowid FROM [{table}] LIMIT 0");
1639            if exec_sql_silent(db, &probe) != SQLITE_OK {
1640                set_error(
1641                    ctx,
1642                    &format!(
1643                        "CreateSpatialIndex: table [{table}] has no rowid column. \
1644                     WITHOUT ROWID tables are not supported. Recreate the table \
1645                     without the WITHOUT ROWID clause, or verify the table exists."
1646                    ),
1647                );
1648                rollback_savepoint(db, ctx, savepoint);
1649                return;
1650            }
1651
1652            // 1. Create the R-tree virtual table if missing.
1653            let sql = format!(
1654            "CREATE VIRTUAL TABLE IF NOT EXISTS [{rtree}] USING rtree(id, xmin, xmax, ymin, ymax)"
1655        );
1656            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1657                rollback_savepoint(db, ctx, savepoint);
1658                return;
1659            }
1660
1661            // 2. Rebuild index contents from the base table (idempotent on repeated calls).
1662            let sql = format!("DELETE FROM [{rtree}]");
1663            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1664                rollback_savepoint(db, ctx, savepoint);
1665                return;
1666            }
1667
1668            let sql = format!(
1669                "INSERT INTO [{rtree}] \
1670             SELECT rowid, ST_XMin([{column}]), ST_XMax([{column}]), \
1671             ST_YMin([{column}]), ST_YMax([{column}]) \
1672             FROM [{table}] WHERE [{column}] IS NOT NULL AND ST_IsEmpty([{column}]) = 0"
1673            );
1674            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1675                rollback_savepoint(db, ctx, savepoint);
1676                return;
1677            }
1678
1679            // 3. AFTER INSERT trigger
1680            let trigger_insert = format!("{table}_{column}_insert");
1681            let sql = format!(
1682                "CREATE TRIGGER IF NOT EXISTS [{trigger_insert}] AFTER INSERT ON [{table}] \
1683             WHEN NEW.[{column}] IS NOT NULL AND ST_IsEmpty(NEW.[{column}]) = 0 \
1684             BEGIN \
1685               INSERT INTO [{rtree}] VALUES ( \
1686                 NEW.rowid, \
1687                 ST_XMin(NEW.[{column}]), ST_XMax(NEW.[{column}]), \
1688                 ST_YMin(NEW.[{column}]), ST_YMax(NEW.[{column}]) \
1689               ); \
1690             END"
1691            );
1692            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1693                rollback_savepoint(db, ctx, savepoint);
1694                return;
1695            }
1696
1697            // 4. AFTER UPDATE trigger. Broad UPDATE so that rowid changes
1698            // (UPDATE ... SET rowid = ... or via INTEGER PRIMARY KEY rewrite)
1699            // still propagate to the index. The WHEN clause skips the DELETE
1700            // plus INSERT when neither the geometry blob nor the rowid changed,
1701            // which is the common case for UPDATEs that only touch unrelated
1702            // columns.
1703            let trigger_update = format!("{table}_{column}_update");
1704            let sql = format!(
1705                "CREATE TRIGGER IF NOT EXISTS [{trigger_update}] AFTER UPDATE ON [{table}] \
1706             WHEN OLD.[{column}] IS NOT NEW.[{column}] OR OLD.rowid IS NOT NEW.rowid \
1707             BEGIN \
1708               DELETE FROM [{rtree}] WHERE id = OLD.rowid; \
1709               INSERT INTO [{rtree}] \
1710                 SELECT NEW.rowid, \
1711                   ST_XMin(NEW.[{column}]), ST_XMax(NEW.[{column}]), \
1712                   ST_YMin(NEW.[{column}]), ST_YMax(NEW.[{column}]) \
1713                 WHERE NEW.[{column}] IS NOT NULL AND ST_IsEmpty(NEW.[{column}]) = 0; \
1714             END"
1715            );
1716            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1717                rollback_savepoint(db, ctx, savepoint);
1718                return;
1719            }
1720
1721            // 5. AFTER DELETE trigger
1722            let trigger_delete = format!("{table}_{column}_delete");
1723            let sql = format!(
1724                "CREATE TRIGGER IF NOT EXISTS [{trigger_delete}] AFTER DELETE ON [{table}] \
1725             BEGIN \
1726               DELETE FROM [{rtree}] WHERE id = OLD.rowid; \
1727             END"
1728            );
1729            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1730                rollback_savepoint(db, ctx, savepoint);
1731                return;
1732            }
1733
1734            let sql = format!(
1735                "INSERT INTO [{SPATIAL_INDEX_CATALOG_TABLE}] (prefix, table_name, column_name) \
1736             VALUES ('{prefix}', '{table}', '{column}') \
1737             ON CONFLICT(prefix) DO UPDATE SET \
1738             table_name = excluded.table_name, \
1739             column_name = excluded.column_name"
1740            );
1741            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1742                rollback_savepoint(db, ctx, savepoint);
1743                return;
1744            }
1745
1746            if exec_sql(db, ctx, &format!("RELEASE {savepoint}")) != SQLITE_OK {
1747                return;
1748            }
1749
1750            set_i32(ctx, 1);
1751        });
1752    }
1753}
1754
1755unsafe extern "C" fn drop_spatial_index_xfunc(
1756    ctx: *mut sqlite3_context,
1757    _n: c_int,
1758    argv: *mut *mut sqlite3_value,
1759) {
1760    unsafe {
1761        xfunc_guard(ctx, "DropSpatialIndex", || {
1762            let Some((table, column)) = get_table_column(ctx, argv, "DropSpatialIndex") else {
1763                return;
1764            };
1765
1766            let db = sqlite3_context_db_handle(ctx);
1767            let prefix = format!("{table}_{column}");
1768            let savepoint = "sqlitegis_drop_spatial_index";
1769
1770            if exec_sql(db, ctx, &format!("SAVEPOINT {savepoint}")) != SQLITE_OK {
1771                return;
1772            }
1773
1774            if !ensure_spatial_index_catalog_table(db, ctx, "DropSpatialIndex") {
1775                rollback_savepoint(db, ctx, savepoint);
1776                return;
1777            }
1778            if !validate_spatial_index_catalog_shape(db, ctx, "DropSpatialIndex") {
1779                rollback_savepoint(db, ctx, savepoint);
1780                return;
1781            }
1782
1783            let ownership = match ensure_spatial_index_objects_owned_by_table(
1784                db,
1785                ctx,
1786                table,
1787                column,
1788                "DropSpatialIndex",
1789            ) {
1790                Some(v) => v,
1791                None => {
1792                    rollback_savepoint(db, ctx, savepoint);
1793                    return;
1794                }
1795            };
1796            if ownership == SpatialIndexOwnership::Absent {
1797                if exec_sql(db, ctx, &format!("RELEASE {savepoint}")) != SQLITE_OK {
1798                    return;
1799                }
1800                set_i32(ctx, 1);
1801                return;
1802            }
1803
1804            // Drop triggers first, then the R-tree table. Ownership has already
1805            // been verified to avoid cross-table collisions on derived names.
1806            for suffix in &["_insert", "_update", "_delete"] {
1807                let sql = format!("DROP TRIGGER IF EXISTS [{prefix}{suffix}]");
1808                if exec_sql(db, ctx, &sql) != SQLITE_OK {
1809                    rollback_savepoint(db, ctx, savepoint);
1810                    return;
1811                }
1812            }
1813            let sql = format!("DROP TABLE IF EXISTS [{prefix}_rtree]");
1814            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1815                rollback_savepoint(db, ctx, savepoint);
1816                return;
1817            }
1818
1819            let sql =
1820                format!("DELETE FROM [{SPATIAL_INDEX_CATALOG_TABLE}] WHERE prefix = '{prefix}'");
1821            if exec_sql(db, ctx, &sql) != SQLITE_OK {
1822                rollback_savepoint(db, ctx, savepoint);
1823                return;
1824            }
1825
1826            if exec_sql(db, ctx, &format!("RELEASE {savepoint}")) != SQLITE_OK {
1827                return;
1828            }
1829
1830            set_i32(ctx, 1);
1831        });
1832    }
1833}
1834
1835// Registration
1836
1837type XFunc = unsafe extern "C" fn(*mut sqlite3_context, c_int, *mut *mut sqlite3_value);
1838
1839#[derive(Clone, Copy)]
1840struct SqliteCallbackSpec {
1841    name: &'static str,
1842    n_arg: i32,
1843    xfunc: XFunc,
1844}
1845
1846macro_rules! callback_spec {
1847    ($name:literal, $n_arg:literal, $xfunc:ident) => {
1848        SqliteCallbackSpec {
1849            name: $name,
1850            n_arg: $n_arg,
1851            xfunc: $xfunc,
1852        }
1853    };
1854}
1855
1856include!("deterministic_callbacks.rs");
1857include!("direct_only_callbacks.rs");
1858
1859const fn const_str_eq(a: &str, b: &str) -> bool {
1860    let a_bytes = a.as_bytes();
1861    let b_bytes = b.as_bytes();
1862    if a_bytes.len() != b_bytes.len() {
1863        return false;
1864    }
1865
1866    let mut i = 0;
1867    while i < a_bytes.len() {
1868        if a_bytes[i] != b_bytes[i] {
1869            return false;
1870        }
1871        i += 1;
1872    }
1873    true
1874}
1875
1876const fn assert_catalog_callback_parity(
1877    catalog: &[SqliteFunctionSpec],
1878    callbacks: &[SqliteCallbackSpec],
1879) {
1880    assert!(catalog.len() == callbacks.len());
1881    let mut i = 0;
1882    while i < callbacks.len() {
1883        assert!(const_str_eq(callbacks[i].name, catalog[i].name));
1884        assert!(callbacks[i].n_arg == catalog[i].n_arg);
1885        i += 1;
1886    }
1887}
1888
1889const _: () = assert_catalog_callback_parity(
1890    SQLITE_DETERMINISTIC_FUNCTIONS,
1891    SQLITE_DETERMINISTIC_CALLBACKS,
1892);
1893const _: () =
1894    assert_catalog_callback_parity(SQLITE_DIRECT_ONLY_FUNCTIONS, SQLITE_DIRECT_ONLY_CALLBACKS);
1895
1896unsafe fn reg(db: *mut sqlite3, name: &str, n_arg: c_int, flags: c_int, xfunc: XFunc) -> c_int {
1897    unsafe {
1898        let c_name = match CString::new(name) {
1899            Ok(v) => v,
1900            Err(_) => return SQLITE_ERROR,
1901        };
1902        sqlite3_create_function_v2(
1903            db,
1904            c_name.as_ptr(),
1905            n_arg,
1906            flags,
1907            std::ptr::null_mut(),
1908            Some(xfunc),
1909            None,
1910            None,
1911            None,
1912        )
1913    }
1914}
1915
1916/// Register all SQLiteGIS spatial functions into an open SQLite database.
1917///
1918/// Returns `SQLITE_OK` (0) on success, or the first error code on failure.
1919///
1920/// # Safety
1921/// `db` must be a valid, open SQLite database handle for the lifetime of the call.
1922///
1923/// # Example
1924///
1925/// Open an in-memory connection via the raw `libsqlite3-sys` FFI, register
1926/// the SQLiteGIS spatial functions, then call one from SQL:
1927///
1928/// ```
1929/// use std::ffi::{CStr, CString};
1930/// use std::ptr;
1931/// use libsqlite3_sys::{
1932///     sqlite3, sqlite3_close, sqlite3_column_text, sqlite3_finalize,
1933///     sqlite3_open, sqlite3_prepare_v2, sqlite3_step,
1934///     SQLITE_OK, SQLITE_ROW,
1935/// };
1936///
1937/// unsafe {
1938///     let mut db: *mut sqlite3 = ptr::null_mut();
1939///     let path = CString::new(":memory:").unwrap();
1940///     assert_eq!(sqlite3_open(path.as_ptr(), &mut db), SQLITE_OK);
1941///
1942///     // Wire SQLiteGIS's spatial functions into this connection.
1943///     assert_eq!(sqlitegis::sqlite::register_functions(db), SQLITE_OK);
1944///
1945///     // ST_AsText, ST_Point, and the rest of the catalogue are now
1946///     // callable from any prepared statement on this handle.
1947///     let sql = CString::new("SELECT ST_AsText(ST_Point(1.0, 2.0, 4326))").unwrap();
1948///     let mut stmt = ptr::null_mut();
1949///     assert_eq!(
1950///         sqlite3_prepare_v2(db, sql.as_ptr(), -1, &mut stmt, ptr::null_mut()),
1951///         SQLITE_OK,
1952///     );
1953///     assert_eq!(sqlite3_step(stmt), SQLITE_ROW);
1954///     let text = CStr::from_ptr(sqlite3_column_text(stmt, 0).cast());
1955///     assert_eq!(text.to_str().unwrap(), "POINT(1 2)");
1956///     sqlite3_finalize(stmt);
1957///     sqlite3_close(db);
1958/// }
1959/// ```
1960pub unsafe fn register_functions(db: *mut sqlite3) -> c_int {
1961    unsafe {
1962        for callback in SQLITE_DETERMINISTIC_CALLBACKS {
1963            let rc = reg(
1964                db,
1965                callback.name,
1966                callback.n_arg as c_int,
1967                DET,
1968                callback.xfunc,
1969            );
1970            if rc != SQLITE_OK {
1971                return rc;
1972            }
1973        }
1974
1975        for callback in SQLITE_DIRECT_ONLY_CALLBACKS {
1976            let rc = reg(
1977                db,
1978                callback.name,
1979                callback.n_arg as c_int,
1980                DIRECT,
1981                callback.xfunc,
1982            );
1983            if rc != SQLITE_OK {
1984                return rc;
1985            }
1986        }
1987
1988        SQLITE_OK
1989    }
1990}
1991
1992/// Register SQLiteGIS as a SQLite auto-extension: from the next call onward,
1993/// every new SQLite connection opened in this process has the SQLiteGIS
1994/// spatial functions registered on it automatically.
1995///
1996/// This is the recommended way to wire SQLiteGIS into [Diesel]'s
1997/// `SqliteConnection::establish(...)` flow, which does not give the caller a
1998/// raw `*mut sqlite3` to pass to [`register_functions`] directly.
1999///
2000/// Calling the function more than once is a no-op: a `std::sync::Once`
2001/// guarantees the underlying `sqlite3_auto_extension` registration happens
2002/// at most once per process.
2003///
2004/// Native targets only. On wasm32 the connection lifecycle is owned by
2005/// `sqlite-wasm-rs` and you should call [`register_functions`] from your
2006/// own `auto_extension` shim once per connection instead.
2007///
2008/// Returns `SQLITE_OK` on the first call, and `SQLITE_OK` on every
2009/// subsequent call too. A non-zero return from this function should be
2010/// treated as a SQLite-internal failure.
2011///
2012/// [Diesel]: https://diesel.rs/
2013///
2014/// ```
2015/// use diesel::{Connection, RunQueryDsl, sqlite::SqliteConnection};
2016/// use diesel::deserialize::QueryableByName;
2017/// use diesel::sql_types::Text;
2018///
2019/// sqlitegis::sqlite::register_on_every_new_connection();
2020/// let mut c = SqliteConnection::establish(":memory:").unwrap();
2021///
2022/// #[derive(QueryableByName)]
2023/// struct R { #[diesel(sql_type = Text)] wkt: String }
2024///
2025/// let row: R = diesel::sql_query(
2026///     "SELECT ST_AsText(ST_Point(1.0, 2.0, 4326)) AS wkt",
2027/// ).get_result(&mut c).unwrap();
2028/// assert_eq!(row.wkt, "POINT(1 2)");
2029/// ```
2030pub fn register_on_every_new_connection() -> c_int {
2031    use std::sync::Once;
2032    static INIT: Once = Once::new();
2033    INIT.call_once(|| unsafe {
2034        sqlite3_auto_extension(Some(sqlitegis_auto_extension_init));
2035    });
2036    SQLITE_OK
2037}
2038
2039unsafe extern "C" fn sqlitegis_auto_extension_init(
2040    db: *mut sqlite3,
2041    _pz_err_msg: *mut *mut std::os::raw::c_char,
2042    _p_api: *const sqlite3_api_routines,
2043) -> c_int {
2044    unsafe { register_functions(db) }
2045}
2046
2047// C entry point for loadable extension (native only)
2048//
2049// Gated on feature = "sqlite-extension" so consumers that enable only the
2050// in-process `sqlite` feature do not silently re-export `sqlite3_sqlitegis_init`
2051// from their own binaries. That would collide with anyone else embedding
2052// the SQLiteGIS extension at the C ABI level.
2053
2054/// `sqlite3_sqlitegis_init` is the entry point called by SQLite when loading
2055/// this library as a loadable extension (`SELECT load_extension('libsqlitegis')`).
2056#[cfg(all(feature = "sqlite-extension", not(target_arch = "wasm32")))]
2057#[no_mangle]
2058pub unsafe extern "C" fn sqlite3_sqlitegis_init(
2059    db: *mut sqlite3,
2060    _pz_err_msg: *mut *mut std::ffi::c_char,
2061    _p_api: *mut sqlite3_api_routines,
2062) -> c_int {
2063    unsafe {
2064        match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| register_functions(db))) {
2065            Ok(rc) => rc,
2066            Err(_) => SQLITE_ERROR,
2067        }
2068    }
2069}
2070
2071#[cfg(test)]
2072mod tests {
2073    use super::*;
2074    use crate::core::function_catalog::{SemanticCase, SemanticExpectation};
2075    use std::ffi::{CStr, CString};
2076    use std::ptr;
2077
2078    unsafe extern "C" fn guarded_constant_xfunc(
2079        ctx: *mut sqlite3_context,
2080        _n: c_int,
2081        _argv: *mut *mut sqlite3_value,
2082    ) {
2083        unsafe {
2084            xfunc_guard(ctx, "GuardedConstant", || {
2085                set_i32(ctx, 7);
2086            });
2087        }
2088    }
2089
2090    unsafe extern "C" fn guarded_panic_xfunc(
2091        ctx: *mut sqlite3_context,
2092        _n: c_int,
2093        _argv: *mut *mut sqlite3_value,
2094    ) {
2095        unsafe {
2096            xfunc_guard(ctx, "GuardedPanic", || {
2097                panic!("boom");
2098            });
2099        }
2100    }
2101
2102    unsafe fn open_db() -> *mut sqlite3 {
2103        unsafe {
2104            let mut db = ptr::null_mut();
2105            let path = CString::new(":memory:").expect("valid sqlite path");
2106            assert_eq!(sqlite3_open(path.as_ptr(), &mut db), SQLITE_OK);
2107            db
2108        }
2109    }
2110
2111    unsafe fn close_db(db: *mut sqlite3) {
2112        unsafe {
2113            assert_eq!(sqlite3_close(db), SQLITE_OK);
2114        }
2115    }
2116
2117    unsafe fn query_i64(db: *mut sqlite3, sql: &str) -> Result<i64, String> {
2118        unsafe {
2119            let sql_c = CString::new(sql).expect("valid SQL");
2120            let mut stmt = ptr::null_mut();
2121            let rc = sqlite3_prepare_v2(db, sql_c.as_ptr(), -1, &mut stmt, ptr::null_mut());
2122            if rc != SQLITE_OK {
2123                let err = CStr::from_ptr(sqlite3_errmsg(db))
2124                    .to_string_lossy()
2125                    .into_owned();
2126                return Err(err);
2127            }
2128
2129            let step = sqlite3_step(stmt);
2130            if step != SQLITE_ROW {
2131                sqlite3_finalize(stmt);
2132                let err = CStr::from_ptr(sqlite3_errmsg(db))
2133                    .to_string_lossy()
2134                    .into_owned();
2135                return Err(err);
2136            }
2137
2138            let value = sqlite3_column_int64(stmt, 0);
2139            sqlite3_finalize(stmt);
2140            Ok(value)
2141        }
2142    }
2143
2144    #[derive(Debug)]
2145    enum QueryValue {
2146        Null,
2147        Integer(i64),
2148        Float(f64),
2149        Text(String),
2150        Blob(Vec<u8>),
2151    }
2152
2153    unsafe fn query_value(db: *mut sqlite3, sql: &str) -> Result<QueryValue, String> {
2154        unsafe {
2155            let sql_c = CString::new(sql).expect("valid SQL");
2156            let mut stmt = ptr::null_mut();
2157            let rc = sqlite3_prepare_v2(db, sql_c.as_ptr(), -1, &mut stmt, ptr::null_mut());
2158            if rc != SQLITE_OK {
2159                let err = CStr::from_ptr(sqlite3_errmsg(db))
2160                    .to_string_lossy()
2161                    .into_owned();
2162                return Err(err);
2163            }
2164
2165            let step = sqlite3_step(stmt);
2166            if step != SQLITE_ROW {
2167                sqlite3_finalize(stmt);
2168                let err = CStr::from_ptr(sqlite3_errmsg(db))
2169                    .to_string_lossy()
2170                    .into_owned();
2171                return Err(err);
2172            }
2173
2174            let value = match sqlite3_column_type(stmt, 0) {
2175                SQLITE_NULL => QueryValue::Null,
2176                SQLITE_INTEGER => QueryValue::Integer(sqlite3_column_int64(stmt, 0)),
2177                SQLITE_FLOAT => QueryValue::Float(sqlite3_column_double(stmt, 0)),
2178                SQLITE_TEXT => {
2179                    let ptr = sqlite3_column_text(stmt, 0);
2180                    if ptr.is_null() {
2181                        sqlite3_finalize(stmt);
2182                        return Err("unexpected NULL text pointer for SQLITE_TEXT".to_string());
2183                    }
2184                    QueryValue::Text(CStr::from_ptr(ptr.cast()).to_string_lossy().into_owned())
2185                }
2186                SQLITE_BLOB => {
2187                    let len = sqlite3_column_bytes(stmt, 0) as usize;
2188                    let ptr = sqlite3_column_blob(stmt, 0) as *const u8;
2189                    if len == 0 {
2190                        QueryValue::Blob(Vec::new())
2191                    } else if ptr.is_null() {
2192                        sqlite3_finalize(stmt);
2193                        return Err("unexpected NULL blob pointer for SQLITE_BLOB".to_string());
2194                    } else {
2195                        QueryValue::Blob(std::slice::from_raw_parts(ptr, len).to_vec())
2196                    }
2197                }
2198                other => {
2199                    sqlite3_finalize(stmt);
2200                    return Err(format!("unexpected SQLite type code: {other}"));
2201                }
2202            };
2203
2204            sqlite3_finalize(stmt);
2205            Ok(value)
2206        }
2207    }
2208
2209    fn assert_semantic_expectation(
2210        spec: &SqliteFunctionSpec,
2211        case: &SemanticCase,
2212        result: Result<QueryValue, String>,
2213    ) {
2214        match case.expected {
2215            SemanticExpectation::Null => match result {
2216                Ok(QueryValue::Null) => {}
2217                Ok(got) => panic!(
2218                    "expected NULL for {}({}) case `{}` via `{}`, got {:?}",
2219                    spec.name, spec.n_arg, case.id, case.sql, got
2220                ),
2221                Err(err) => panic!(
2222                    "expected NULL for {}({}) case `{}` via `{}`, got error: {err}",
2223                    spec.name, spec.n_arg, case.id, case.sql
2224                ),
2225            },
2226            SemanticExpectation::NumericFinite => match result {
2227                Ok(QueryValue::Integer(_)) => {}
2228                Ok(QueryValue::Float(v)) => {
2229                    assert!(
2230                        v.is_finite(),
2231                        "non-finite numeric for {}({}) case `{}` via `{}`",
2232                        spec.name,
2233                        spec.n_arg,
2234                        case.id,
2235                        case.sql
2236                    );
2237                }
2238                Ok(got) => panic!(
2239                    "expected numeric for {}({}) case `{}` via `{}`, got {:?}",
2240                    spec.name, spec.n_arg, case.id, case.sql, got
2241                ),
2242                Err(err) => panic!(
2243                    "expected numeric for {}({}) case `{}` via `{}`, got error: {err}",
2244                    spec.name, spec.n_arg, case.id, case.sql
2245                ),
2246            },
2247            SemanticExpectation::TextNonEmpty => match result {
2248                Ok(QueryValue::Text(v)) => {
2249                    assert!(
2250                        !v.is_empty(),
2251                        "empty text result for {}({}) case `{}` via `{}`",
2252                        spec.name,
2253                        spec.n_arg,
2254                        case.id,
2255                        case.sql
2256                    );
2257                }
2258                Ok(got) => panic!(
2259                    "expected non-empty text for {}({}) case `{}` via `{}`, got {:?}",
2260                    spec.name, spec.n_arg, case.id, case.sql, got
2261                ),
2262                Err(err) => panic!(
2263                    "expected non-empty text for {}({}) case `{}` via `{}`, got error: {err}",
2264                    spec.name, spec.n_arg, case.id, case.sql
2265                ),
2266            },
2267            SemanticExpectation::BlobNonEmpty => match result {
2268                Ok(QueryValue::Blob(v)) => {
2269                    assert!(
2270                        !v.is_empty(),
2271                        "empty blob result for {}({}) case `{}` via `{}`",
2272                        spec.name,
2273                        spec.n_arg,
2274                        case.id,
2275                        case.sql
2276                    );
2277                }
2278                Ok(got) => panic!(
2279                    "expected non-empty blob for {}({}) case `{}` via `{}`, got {:?}",
2280                    spec.name, spec.n_arg, case.id, case.sql, got
2281                ),
2282                Err(err) => panic!(
2283                    "expected non-empty blob for {}({}) case `{}` via `{}`, got error: {err}",
2284                    spec.name, spec.n_arg, case.id, case.sql
2285                ),
2286            },
2287            SemanticExpectation::Bool01 => match result {
2288                Ok(QueryValue::Integer(v)) => {
2289                    assert!(
2290                        v == 0 || v == 1,
2291                        "bool result must be 0/1 for {}({}) case `{}` via `{}`, got {v}",
2292                        spec.name,
2293                        spec.n_arg,
2294                        case.id,
2295                        case.sql
2296                    );
2297                }
2298                Ok(got) => panic!(
2299                    "expected bool-as-int for {}({}) case `{}` via `{}`, got {:?}",
2300                    spec.name, spec.n_arg, case.id, case.sql, got
2301                ),
2302                Err(err) => panic!(
2303                    "expected bool-as-int for {}({}) case `{}` via `{}`, got error: {err}",
2304                    spec.name, spec.n_arg, case.id, case.sql
2305                ),
2306            },
2307            SemanticExpectation::ErrorContains(expected_substring) => match result {
2308                Err(err) => assert!(
2309                    err.contains(expected_substring),
2310                    "error mismatch for {}({}) case `{}` via `{}`: expected substring `{}`, got `{}`",
2311                    spec.name,
2312                    spec.n_arg,
2313                    case.id,
2314                    case.sql,
2315                    expected_substring,
2316                    err
2317                ),
2318                Ok(got) => panic!(
2319                    "expected SQL error for {}({}) case `{}` via `{}`, got {:?}",
2320                    spec.name, spec.n_arg, case.id, case.sql, got
2321                ),
2322            },
2323        }
2324    }
2325
2326    #[test]
2327    fn checked_c_int_len_accepts_small_and_boundary_values() {
2328        assert_eq!(checked_c_int_len(0), Some(0));
2329        assert_eq!(checked_c_int_len(1), Some(1));
2330        assert_eq!(checked_c_int_len(c_int::MAX as usize), Some(c_int::MAX));
2331    }
2332
2333    #[test]
2334    fn checked_c_int_len_rejects_values_larger_than_c_int() {
2335        assert_eq!(checked_c_int_len((c_int::MAX as usize) + 1), None);
2336        assert_eq!(checked_c_int_len(usize::MAX), None);
2337    }
2338
2339    #[test]
2340    fn sql_to_cstring_accepts_sql_without_nul() {
2341        let c_sql = sql_to_cstring("SELECT 1").expect("valid SQL should convert to CString");
2342        assert_eq!(c_sql.as_c_str().to_bytes(), b"SELECT 1");
2343    }
2344
2345    #[test]
2346    fn sql_to_cstring_rejects_sql_with_nul() {
2347        assert!(sql_to_cstring("SELECT\0 1").is_err());
2348    }
2349
2350    #[test]
2351    fn xfunc_guard_allows_normal_execution() {
2352        unsafe {
2353            let db = open_db();
2354
2355            let func_name = CString::new("GuardedConstant").expect("valid function name");
2356            let rc = sqlite3_create_function_v2(
2357                db,
2358                func_name.as_ptr(),
2359                0,
2360                SQLITE_UTF8,
2361                ptr::null_mut(),
2362                Some(guarded_constant_xfunc),
2363                None,
2364                None,
2365                None,
2366            );
2367            assert_eq!(rc, SQLITE_OK, "function registration should succeed");
2368
2369            let value = query_i64(db, "SELECT GuardedConstant()").expect("query should succeed");
2370            assert_eq!(value, 7);
2371
2372            close_db(db);
2373        }
2374    }
2375
2376    #[test]
2377    fn xfunc_guard_converts_panic_into_sqlite_error() {
2378        unsafe {
2379            let db = open_db();
2380
2381            let func_name = CString::new("GuardedPanic").expect("valid function name");
2382            let rc = sqlite3_create_function_v2(
2383                db,
2384                func_name.as_ptr(),
2385                0,
2386                SQLITE_UTF8,
2387                ptr::null_mut(),
2388                Some(guarded_panic_xfunc),
2389                None,
2390                None,
2391                None,
2392            );
2393            assert_eq!(rc, SQLITE_OK, "function registration should succeed");
2394
2395            let err = query_i64(db, "SELECT GuardedPanic()")
2396                .expect_err("panic should be surfaced as SQL error");
2397            assert!(
2398                err.contains("panic in SQLite callback"),
2399                "unexpected error message: {err}"
2400            );
2401
2402            close_db(db);
2403        }
2404    }
2405
2406    #[test]
2407    fn register_functions_semantic_smoke_covers_full_catalog() {
2408        unsafe {
2409            let db = open_db();
2410
2411            let rc = register_functions(db);
2412            assert_eq!(rc, SQLITE_OK, "register_functions should succeed");
2413
2414            for spec in SQLITE_DETERMINISTIC_FUNCTIONS {
2415                for case in spec.semantic_cases {
2416                    let result = query_value(db, case.sql);
2417                    assert_semantic_expectation(spec, case, result);
2418                }
2419            }
2420
2421            close_db(db);
2422        }
2423    }
2424
2425    #[test]
2426    fn register_functions_direct_only_semantic_smoke() {
2427        unsafe {
2428            let db = open_db();
2429
2430            let rc = register_functions(db);
2431            assert_eq!(rc, SQLITE_OK, "register_functions should succeed");
2432
2433            assert_eq!(
2434                exec_sql_inner(db, "CREATE TABLE _rt(geom BLOB)", None),
2435                SQLITE_OK
2436            );
2437
2438            for spec in SQLITE_DIRECT_ONLY_FUNCTIONS {
2439                for case in spec.semantic_cases {
2440                    let result = query_value(db, case.sql);
2441                    assert_semantic_expectation(spec, case, result);
2442                }
2443            }
2444
2445            close_db(db);
2446        }
2447    }
2448}