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}