Skip to main content

better_duck_core/udf/
replacement.rs

1//! DuckDB replacement scans: automatically rewrite an unresolved table
2//! reference into a table function call — e.g. routing
3//! `SELECT * FROM 'data.parquet'` to `read_parquet('data.parquet')` by file
4//! extension, without the caller writing the function call explicitly.
5//!
6//! # Experimental
7//!
8//! No comparable safe-Rust design exists in the reference `duckdb` crate to
9//! model this on — this is a from-scratch design. It deliberately restricts
10//! what a replacement scan can do: it may only rewrite the unresolved name
11//! into a table function call with literal parameters
12//! ([`ReplacementScanInfo::set_function_name`]/[`add_parameter`](ReplacementScanInfo::add_parameter)).
13//! There is **no way to run SQL from inside the callback** — DuckDB invokes
14//! it while resolving the very query being planned, on the same connection;
15//! issuing another query from there would re-enter (and likely deadlock or
16//! corrupt) that in-progress resolution. If a rewrite needs data that can
17//! only come from a query, compute it before the query that triggers the
18//! scan, not inside the callback.
19
20use std::ffi::{c_void, CStr, CString};
21
22use crate::{
23    database::Database,
24    error::{Error, Result},
25    ffi::{
26        duckdb_add_replacement_scan, duckdb_destroy_value, duckdb_replacement_scan_add_parameter,
27        duckdb_replacement_scan_info, duckdb_replacement_scan_set_error,
28        duckdb_replacement_scan_set_function_name,
29    },
30    types::DuckDialect,
31};
32
33use super::callback::{contain_callback, CallbackErrorSink};
34
35/// A hook that rewrites an unresolved table reference into a table function
36/// call. See the module docs for the v1 scope restriction (no SQL execution).
37///
38/// Stateless by design, matching [`VTab`](super::VTab)/[`VScalar`](super::VScalar):
39/// `Self` is never constructed, only used as a marker type for
40/// [`Database::register_replacement_scan`].
41pub trait ReplacementScan {
42    /// Called whenever DuckDB can't resolve `table_name` to an existing
43    /// table, view, or CTE. Call [`info.set_function_name(...)`](ReplacementScanInfo::set_function_name)
44    /// (optionally followed by [`info.add_parameter(...)`](ReplacementScanInfo::add_parameter))
45    /// to rewrite it into a table function call; otherwise leave `info`
46    /// untouched to decline, and DuckDB reports its normal "table not found"
47    /// error.
48    ///
49    /// # Errors
50    ///
51    /// Returns an error to fail the query with that message instead of
52    /// DuckDB's default "table not found".
53    fn replace(
54        table_name: &str,
55        info: &ReplacementScanInfo,
56    ) -> super::UdfResult<()>;
57}
58
59/// An interface to redirect an unresolved table reference during a
60/// replacement scan callback.
61pub struct ReplacementScanInfo {
62    ptr: duckdb_replacement_scan_info,
63}
64
65impl ReplacementScanInfo {
66    fn from(ptr: duckdb_replacement_scan_info) -> Self {
67        Self { ptr }
68    }
69
70    /// Rewrites the unresolved table reference into a call to the table
71    /// function named `function_name`.
72    ///
73    /// Must be called for the rewrite to take effect at all — if
74    /// [`ReplacementScan::replace`] returns `Ok(())` without ever calling
75    /// this, DuckDB reports its normal "table not found" error, as if no
76    /// replacement scan had run.
77    ///
78    /// # Errors
79    ///
80    /// Returns an error if `function_name` contains a NUL byte.
81    pub fn set_function_name(
82        &self,
83        function_name: &str,
84    ) -> Result<()> {
85        let c_name = CString::new(function_name)?;
86        // SAFETY: `self.ptr` is valid for the duration of the callback;
87        // `c_name` is a valid, NUL-terminated C string for the duration of
88        // this call.
89        unsafe { duckdb_replacement_scan_set_function_name(self.ptr, c_name.as_ptr()) };
90        Ok(())
91    }
92
93    /// Adds a literal parameter to the table function call being built, in
94    /// the order added.
95    ///
96    /// There is deliberately no way to bind a *computed* parameter that would
97    /// require running SQL — see the module docs.
98    ///
99    /// # Errors
100    ///
101    /// Returns an error if `value` cannot be converted to a DuckDB value.
102    pub fn add_parameter<T: DuckDialect>(
103        &self,
104        value: &T,
105    ) -> Result<()> {
106        let mut v = value.to_duck().map_err(Error::ConversionError)?;
107        // SAFETY: `self.ptr` is valid; `v` was just created above. This
108        // function does not take ownership of `v` (matching
109        // `duckdb_bind_value`/`duckdb_append_value`'s convention elsewhere in
110        // this crate) — destroyed exactly once below.
111        unsafe { duckdb_replacement_scan_add_parameter(self.ptr, v) };
112        // SAFETY: `v` was created above and not yet destroyed.
113        unsafe { duckdb_destroy_value(&mut v) };
114        Ok(())
115    }
116}
117
118impl CallbackErrorSink for ReplacementScanInfo {
119    fn set_c_error(
120        &self,
121        error: &CStr,
122    ) {
123        // SAFETY: `self.ptr` is valid for the duration of the callback;
124        // `error` is a valid, NUL-terminated C string.
125        unsafe { duckdb_replacement_scan_set_error(self.ptr, error.as_ptr()) };
126    }
127}
128
129/// The C trampoline installed via `duckdb_add_replacement_scan`.
130///
131/// See [`scalar_trampoline`](super::scalar) for why containment must be the
132/// outermost thing here.
133unsafe extern "C" fn replacement_scan_trampoline<T: ReplacementScan>(
134    info: duckdb_replacement_scan_info,
135    table_name: *const std::os::raw::c_char,
136    _extra_data: *mut c_void,
137) {
138    let info = ReplacementScanInfo::from(info);
139    contain_callback(&info, || {
140        // SAFETY: DuckDB always passes a valid, NUL-terminated table name for
141        // the duration of this call.
142        let name = unsafe { CStr::from_ptr(table_name) }.to_string_lossy();
143        T::replace(&name, &info)
144    });
145}
146
147impl Database {
148    /// Registers a replacement scan: `T::replace` is called whenever DuckDB
149    /// can't resolve a table reference to an existing table, view, or CTE, on
150    /// every connection to this database (replacement scans are scoped
151    /// per-database, not per-connection).
152    ///
153    /// `T` carries no instance data — like [`VTab`](super::VTab)/
154    /// [`VScalar`](super::VScalar), it's used purely as a marker type; the
155    /// DuckDB C API itself has no failure mode for this registration.
156    pub fn register_replacement_scan<T: ReplacementScan>(&self) {
157        // SAFETY: `self.raw_db()` is a valid, open duckdb_database.
158        // `replacement_scan_trampoline::<T>` matches
159        // `duckdb_replacement_callback_t`'s signature. No extra data is
160        // passed, so no delete callback is needed.
161        unsafe {
162            duckdb_add_replacement_scan(
163                self.raw_db(),
164                Some(replacement_scan_trampoline::<T>),
165                std::ptr::null_mut(),
166                None,
167            )
168        };
169    }
170}
171
172#[cfg(test)]
173mod tests {
174    use super::*;
175    use crate::types::value::DuckValue;
176
177    /// Routes any unresolved table name ending in `.range` to `range(n)`,
178    /// where `n` is parsed from the part before the extension — e.g.
179    /// `'5.range'` becomes `range(5)`.
180    struct RangeByExtension;
181
182    impl ReplacementScan for RangeByExtension {
183        fn replace(
184            table_name: &str,
185            info: &ReplacementScanInfo,
186        ) -> super::super::UdfResult<()> {
187            let Some(stem) = table_name.strip_suffix(".range") else {
188                return Ok(()); // Decline: not our extension.
189            };
190            let n: i64 =
191                stem.parse().map_err(|_| format!("'{table_name}' has a non-integer stem"))?;
192            info.set_function_name("range")?;
193            info.add_parameter(&n)?;
194            Ok(())
195        }
196    }
197
198    #[test]
199    fn replacement_scan_rewrites_unresolved_table_reference() {
200        let db = crate::database::Database::open_in_memory().unwrap();
201        db.register_replacement_scan::<RangeByExtension>();
202        let mut conn = db.connect().unwrap();
203
204        let result = conn.execute("SELECT * FROM '5.range' ORDER BY range").unwrap();
205        let rows: Vec<_> = result.collect::<Result<_>>().unwrap();
206        assert_eq!(rows.len(), 5);
207    }
208
209    #[test]
210    fn replacement_scan_declining_leaves_the_normal_error() {
211        let db = crate::database::Database::open_in_memory().unwrap();
212        db.register_replacement_scan::<RangeByExtension>();
213        let mut conn = db.connect().unwrap();
214
215        let err = match conn.execute("SELECT * FROM does_not_exist") {
216            Ok(_) => panic!("expected an error"),
217            Err(e) => e,
218        };
219        assert!(err.to_string().to_lowercase().contains("does_not_exist"), "{err}");
220    }
221
222    #[test]
223    fn replacement_scan_error_surfaces_as_query_error_and_connection_stays_usable() {
224        let db = crate::database::Database::open_in_memory().unwrap();
225        db.register_replacement_scan::<RangeByExtension>();
226        let mut conn = db.connect().unwrap();
227
228        let err = match conn.execute("SELECT * FROM 'not-a-number.range'") {
229            Ok(_) => panic!("expected an error"),
230            Err(e) => e,
231        };
232        assert!(err.to_string().contains("non-integer stem"), "{err}");
233        conn.execute_batch("CREATE TABLE t (v INTEGER)").unwrap();
234    }
235
236    #[test]
237    fn replacement_scan_applies_to_every_connection_on_the_shared_database() {
238        let db = crate::database::Database::open_in_memory().unwrap();
239        db.register_replacement_scan::<RangeByExtension>();
240        let mut a = db.connect().unwrap();
241        let mut b = db.connect().unwrap();
242
243        for conn in [&mut a, &mut b] {
244            let result = conn.execute("SELECT * FROM '3.range' ORDER BY range").unwrap();
245            let rows: Vec<_> = result.collect::<Result<_>>().unwrap();
246            let got: Vec<i64> = rows
247                .iter()
248                .map(|r| match r.get("range").unwrap() {
249                    DuckValue::BigInt(n) => *n,
250                    other => panic!("expected BigInt, got {other:?}"),
251                })
252                .collect();
253            assert_eq!(got, vec![0, 1, 2]);
254        }
255    }
256}