Skip to main content

diesel/sqlite/connection/
functions.rs

1#[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
2extern crate libsqlite3_sys as ffi;
3
4#[cfg(all(target_family = "wasm", target_os = "unknown"))]
5use sqlite_wasm_rs as ffi;
6
7use super::raw::RawConnection;
8use super::{Sqlite, SqliteAggregateFunction, SqliteBindValue, SqliteConnection};
9use crate::backend::Backend;
10use crate::deserialize::{FromSqlRow, StaticallySizedRow};
11use crate::result::{DatabaseErrorKind, Error, QueryResult};
12use crate::row::{Field, PartialRow, Row, RowIndex, RowSealed};
13use crate::serialize::{IsNull, Output, ToSql};
14use crate::sql_types::HasSqlType;
15use crate::sqlite::SqliteFunctionBehavior;
16use crate::sqlite::SqliteValue;
17use crate::sqlite::connection::bind_collector::SqliteBindValueRef;
18use crate::sqlite::connection::sqlite_value::OwnedSqliteValue;
19use alloc::boxed::Box;
20use alloc::string::ToString;
21
22pub(super) fn register<ArgsSqlType, RetSqlType, Args, Ret, F>(
23    conn: &RawConnection,
24    fn_name: &str,
25    behavior: SqliteFunctionBehavior,
26    mut f: F,
27) -> QueryResult<()>
28where
29    F: FnMut(&RawConnection, Args) -> Ret + core::panic::UnwindSafe + Send + 'static,
30    Args: FromSqlRow<ArgsSqlType, Sqlite> + StaticallySizedRow<ArgsSqlType, Sqlite>,
31    Ret: ToSql<RetSqlType, Sqlite>,
32    Sqlite: HasSqlType<RetSqlType>,
33{
34    let fields_needed = Args::FIELD_COUNT;
35    if fields_needed > 127 {
36        return Err(Error::DatabaseError(
37            DatabaseErrorKind::UnableToSendCommand,
38            Box::new("SQLite functions cannot take more than 127 parameters".to_string()),
39        ));
40    }
41
42    conn.register_sql_function(fn_name, fields_needed, behavior, move |conn, args| {
43        let args = build_sql_function_args::<ArgsSqlType, Args>(args, conn.internal_connection)?;
44
45        Ok(f(conn, args))
46    })?;
47    Ok(())
48}
49
50pub(super) fn register_noargs<RetSqlType, Ret, F>(
51    conn: &RawConnection,
52    fn_name: &str,
53    behavior: SqliteFunctionBehavior,
54    mut f: F,
55) -> QueryResult<()>
56where
57    F: FnMut() -> Ret + core::panic::UnwindSafe + Send + 'static,
58    Ret: ToSql<RetSqlType, Sqlite>,
59    Sqlite: HasSqlType<RetSqlType>,
60{
61    conn.register_sql_function(fn_name, 0, behavior, move |_, _| Ok(f()))?;
62    Ok(())
63}
64
65pub(super) fn register_aggregate<ArgsSqlType, RetSqlType, Args, Ret, A>(
66    conn: &RawConnection,
67    fn_name: &str,
68    behavior: SqliteFunctionBehavior,
69) -> QueryResult<()>
70where
71    A: SqliteAggregateFunction<Args, Output = Ret> + 'static + Send + core::panic::UnwindSafe,
72    Args: FromSqlRow<ArgsSqlType, Sqlite> + StaticallySizedRow<ArgsSqlType, Sqlite>,
73    Ret: ToSql<RetSqlType, Sqlite>,
74    Sqlite: HasSqlType<RetSqlType>,
75{
76    let fields_needed = Args::FIELD_COUNT;
77    if fields_needed > 127 {
78        return Err(Error::DatabaseError(
79            DatabaseErrorKind::UnableToSendCommand,
80            Box::new("SQLite functions cannot take more than 127 parameters".to_string()),
81        ));
82    }
83
84    conn.register_aggregate_function::<ArgsSqlType, RetSqlType, Args, Ret, A>(
85        fn_name,
86        fields_needed,
87        behavior,
88    )?;
89
90    Ok(())
91}
92
93pub(super) fn build_sql_function_args<ArgsSqlType, Args>(
94    args: &mut [*mut ffi::sqlite3_value],
95    connection: core::ptr::NonNull<ffi::sqlite3>,
96) -> Result<Args, Error>
97where
98    Args: FromSqlRow<ArgsSqlType, Sqlite>,
99{
100    let row = FunctionRow::new(args, connection);
101    Args::build_from_row(&row).map_err(Error::DeserializationError)
102}
103
104// clippy is wrong here, the let binding is required
105// for lifetime reasons
106#[allow(clippy::let_unit_value)]
107pub(super) fn process_sql_function_result<RetSqlType, Ret>(
108    result: &'_ Ret,
109) -> QueryResult<SqliteBindValueRef<'_>>
110where
111    Ret: ToSql<RetSqlType, Sqlite>,
112    Sqlite: HasSqlType<RetSqlType>,
113{
114    let mut metadata_lookup = ();
115    let value = SqliteBindValue {
116        inner: SqliteBindValueRef::Null,
117    };
118    let mut buf = Output::new(value, &mut metadata_lookup);
119    let is_null = result.to_sql(&mut buf).map_err(Error::SerializationError)?;
120
121    if let IsNull::Yes = is_null {
122        Ok(SqliteBindValueRef::Null)
123    } else {
124        Ok(buf.into_inner().inner)
125    }
126}
127
128struct FunctionRow<'a> {
129    args: &'a [Option<OwnedSqliteValue>],
130    field_count: usize,
131    connection: core::ptr::NonNull<ffi::sqlite3>,
132}
133
134impl FunctionRow<'_> {
135    #[allow(unsafe_code)] // complicated ptr cast
136    fn new(
137        args: &mut [*mut ffi::sqlite3_value],
138        connection: core::ptr::NonNull<ffi::sqlite3>,
139    ) -> Self {
140        let lengths = args.len();
141        let args = unsafe {
142            core::slice::from_raw_parts(
143                // This cast is safe because:
144                // * Casting from a pointer to an array to a pointer to the first array
145                // element is safe
146                // * Casting from a raw pointer to `NonNull<T>` is safe,
147                // because `NonNull` is #[repr(transparent)]
148                // * Casting from `NonNull<T>` to `OwnedSqliteValue` is safe,
149                // as the struct is `#[repr(transparent)]
150                // * Casting from `NonNull<T>` to `Option<NonNull<T>>` as the documentation
151                // states: "This is so that enums may use this forbidden value as a discriminant –
152                // Option<NonNull<T>> has the same size as *mut T"
153                // * The last point remains true for `OwnedSqliteValue` as `#[repr(transparent)]
154                // guarantees the same layout as the inner type
155                args as *mut [*mut ffi::sqlite3_value] as *mut ffi::sqlite3_value
156                    as *mut Option<OwnedSqliteValue>,
157                lengths,
158            )
159        };
160
161        Self {
162            field_count: lengths,
163            args,
164            connection,
165        }
166    }
167}
168
169impl RowSealed for FunctionRow<'_> {}
170
171impl<'a> Row<'a, Sqlite> for FunctionRow<'a> {
172    type Field<'f>
173        = FunctionArgument<'f>
174    where
175        'a: 'f,
176        Self: 'f;
177    type InnerPartialRow = Self;
178
179    fn field_count(&self) -> usize {
180        self.field_count
181    }
182
183    fn get<'b, I>(&'b self, idx: I) -> Option<Self::Field<'b>>
184    where
185        'a: 'b,
186        Self: crate::row::RowIndex<I>,
187    {
188        let col_idx = self.idx(idx)?;
189        Some(FunctionArgument {
190            args: self.args,
191            col_idx,
192            connection: self.connection,
193        })
194    }
195
196    fn partial_row(&self, range: core::ops::Range<usize>) -> PartialRow<'_, Self::InnerPartialRow> {
197        PartialRow::new(self, range)
198    }
199}
200
201impl RowIndex<usize> for FunctionRow<'_> {
202    fn idx(&self, idx: usize) -> Option<usize> {
203        if idx < self.field_count() {
204            Some(idx)
205        } else {
206            None
207        }
208    }
209}
210
211impl<'a> RowIndex<&'a str> for FunctionRow<'_> {
212    fn idx(&self, _idx: &'a str) -> Option<usize> {
213        None
214    }
215}
216
217struct FunctionArgument<'a> {
218    args: &'a [Option<OwnedSqliteValue>],
219    col_idx: usize,
220    connection: core::ptr::NonNull<ffi::sqlite3>,
221}
222
223impl<'a> Field<'a, Sqlite> for FunctionArgument<'a> {
224    fn field_name(&self) -> Option<&str> {
225        None
226    }
227
228    fn is_null(&self) -> bool {
229        self.value().is_none()
230    }
231
232    fn value(&self) -> Option<<Sqlite as Backend>::RawValue<'_>> {
233        SqliteValue::from_function_row(self.args, self.col_idx, self.connection)
234    }
235}
236
237impl SqliteConnection {
238    #[doc(hidden)]
239    pub fn register_sql_function<ArgsSqlType, RetSqlType, Args, Ret, F>(
240        &mut self,
241        fn_name: &str,
242        behavior: SqliteFunctionBehavior,
243        mut f: F,
244    ) -> QueryResult<()>
245    where
246        F: FnMut(Args) -> Ret + core::panic::UnwindSafe + Send + 'static,
247        Args: FromSqlRow<ArgsSqlType, Sqlite> + StaticallySizedRow<ArgsSqlType, Sqlite>,
248        Ret: ToSql<RetSqlType, Sqlite>,
249        Sqlite: HasSqlType<RetSqlType>,
250    {
251        register(&self.raw_connection, fn_name, behavior, move |_, args| {
252            f(args)
253        })
254    }
255
256    #[doc(hidden)]
257    pub fn register_noarg_sql_function<RetSqlType, Ret, F>(
258        &mut self,
259        fn_name: &str,
260        behavior: SqliteFunctionBehavior,
261        f: F,
262    ) -> QueryResult<()>
263    where
264        F: FnMut() -> Ret + core::panic::UnwindSafe + Send + 'static,
265        Ret: ToSql<RetSqlType, Sqlite>,
266        Sqlite: HasSqlType<RetSqlType>,
267    {
268        register_noargs(&self.raw_connection, fn_name, behavior, f)
269    }
270
271    #[doc(hidden)]
272    pub fn register_aggregate_function<ArgsSqlType, RetSqlType, Args, Ret, A>(
273        &mut self,
274        fn_name: &str,
275        behavior: SqliteFunctionBehavior,
276    ) -> QueryResult<()>
277    where
278        A: SqliteAggregateFunction<Args, Output = Ret> + 'static + Send + core::panic::UnwindSafe,
279        Args: FromSqlRow<ArgsSqlType, Sqlite> + StaticallySizedRow<ArgsSqlType, Sqlite>,
280        Ret: ToSql<RetSqlType, Sqlite>,
281        Sqlite: HasSqlType<RetSqlType>,
282    {
283        register_aggregate::<_, _, _, _, A>(&self.raw_connection, fn_name, behavior)
284    }
285
286    /// Register a collation function.
287    ///
288    /// `collation` must always return the same answer given the same inputs.
289    /// If `collation` panics and unwinds the stack, the process is aborted, since it is used
290    /// across a C FFI boundary, which cannot be unwound across and there is no way to
291    /// signal failures via the SQLite interface in this case..
292    ///
293    /// If the name is already registered it will be overwritten.
294    ///
295    /// This method will return an error if registering the function fails, either due to an
296    /// out-of-memory situation or because a collation with that name already exists and is
297    /// currently being used in parallel by a query.
298    ///
299    /// The collation needs to be specified when creating a table:
300    /// `CREATE TABLE my_table ( str TEXT COLLATE MY_COLLATION )`,
301    /// where `MY_COLLATION` corresponds to name passed as `collation_name`.
302    ///
303    /// # Example
304    ///
305    /// ```rust
306    /// # include!("../../doctest_setup.rs");
307    /// #
308    /// # fn main() {
309    /// #     run_test().unwrap();
310    /// # }
311    /// #
312    /// # fn run_test() -> QueryResult<()> {
313    /// #     let mut conn = SqliteConnection::establish(":memory:").unwrap();
314    /// // sqlite NOCASE only works for ASCII characters,
315    /// // this collation allows handling UTF-8 (barring locale differences)
316    /// conn.register_collation("RUSTNOCASE", |rhs, lhs| {
317    ///     rhs.to_lowercase().cmp(&lhs.to_lowercase())
318    /// })
319    /// # }
320    /// ```
321    pub fn register_collation<F>(&mut self, collation_name: &str, collation: F) -> QueryResult<()>
322    where
323        F: Fn(&str, &str) -> core::cmp::Ordering + Send + 'static + core::panic::UnwindSafe,
324    {
325        self.raw_connection
326            .register_collation_function(collation_name, collation)
327    }
328}
329
330// miri doesn't support callback over ffi at this point
331#[cfg(all(test, not(miri)))]
332mod tests {
333    use super::*;
334    use crate::connection::SimpleConnection;
335    use crate::prelude::*;
336    use crate::sql_types::{Integer, Text};
337
338    fn connection() -> SqliteConnection {
339        SqliteConnection::establish(":memory:").unwrap()
340    }
341
342    #[declare_sql_function]
343    extern "SQL" {
344        fn fun_case(x: Text) -> Text;
345        fn my_add(x: Integer, y: Integer) -> Integer;
346        fn answer() -> Integer;
347        fn add_counter(x: Integer) -> Integer;
348
349        #[aggregate]
350        fn my_sum(expr: Integer) -> Integer;
351        #[aggregate]
352        fn range_max(expr1: Integer, expr2: Integer, expr3: Integer) -> Nullable<Integer>;
353    }
354
355    #[diesel_test_helper::test]
356    fn register_custom_function() {
357        let connection = &mut connection();
358        fun_case_utils::register_impl(connection, |x: String| {
359            x.chars()
360                .enumerate()
361                .map(|(i, c)| {
362                    if i % 2 == 0 {
363                        c.to_lowercase().to_string()
364                    } else {
365                        c.to_uppercase().to_string()
366                    }
367                })
368                .collect::<String>()
369        })
370        .unwrap();
371
372        let mapped_string = crate::select(fun_case("foobar"))
373            .get_result::<String>(connection)
374            .unwrap();
375        assert_eq!("fOoBaR", mapped_string);
376    }
377
378    #[diesel_test_helper::test]
379    fn register_multiarg_function() {
380        let connection = &mut connection();
381        my_add_utils::register_impl(connection, |x: i32, y: i32| x + y).unwrap();
382
383        let added = crate::select(my_add(1, 2)).get_result::<i32>(connection);
384        assert_eq!(Ok(3), added);
385    }
386
387    #[diesel_test_helper::test]
388    fn register_noarg_function() {
389        let connection = &mut connection();
390        answer_utils::register_impl(connection, || 42).unwrap();
391
392        let answer = crate::select(answer()).get_result::<i32>(connection);
393        assert_eq!(Ok(42), answer);
394    }
395
396    #[diesel_test_helper::test]
397    fn register_nondeterministic_noarg_function() {
398        let connection = &mut connection();
399        answer_utils::register_nondeterministic_impl(connection, || 42).unwrap();
400
401        let answer = crate::select(answer()).get_result::<i32>(connection);
402        assert_eq!(Ok(42), answer);
403    }
404
405    #[diesel_test_helper::test]
406    fn register_nondeterministic_function() {
407        let connection = &mut connection();
408        let mut y = 0;
409        add_counter_utils::register_nondeterministic_impl(connection, move |x: i32| {
410            y += 1;
411            x + y
412        })
413        .unwrap();
414
415        let added = crate::select((add_counter(1), add_counter(1), add_counter(1)))
416            .get_result::<(i32, i32, i32)>(connection);
417        assert_eq!(Ok((2, 3, 4)), added);
418    }
419
420    #[derive(Default)]
421    struct MySum {
422        sum: i32,
423    }
424
425    impl SqliteAggregateFunction<i32> for MySum {
426        type Output = i32;
427
428        fn step(&mut self, expr: i32) {
429            self.sum += expr;
430        }
431
432        fn finalize(aggregator: Option<Self>) -> Self::Output {
433            aggregator.map(|a| a.sum).unwrap_or_default()
434        }
435    }
436
437    table! {
438        my_sum_example {
439            id -> Integer,
440            value -> Integer,
441        }
442    }
443
444    #[diesel_test_helper::test]
445    fn register_aggregate_function() {
446        use self::my_sum_example::dsl::*;
447
448        let connection = &mut connection();
449        crate::sql_query(
450            "CREATE TABLE my_sum_example (id integer primary key autoincrement, value integer)",
451        )
452        .execute(connection)
453        .unwrap();
454        crate::sql_query("INSERT INTO my_sum_example (value) VALUES (1), (2), (3)")
455            .execute(connection)
456            .unwrap();
457
458        my_sum_utils::register_impl_with_behavior::<MySum, _>(
459            connection,
460            SqliteFunctionBehavior::DETERMINISTIC,
461        )
462        .unwrap();
463
464        let result = my_sum_example
465            .select(my_sum(value))
466            .get_result::<i32>(connection);
467        assert_eq!(Ok(6), result);
468    }
469
470    #[diesel_test_helper::test]
471    fn register_aggregate_function_returns_finalize_default_on_empty_set() {
472        use self::my_sum_example::dsl::*;
473
474        let connection = &mut connection();
475        crate::sql_query(
476            "CREATE TABLE my_sum_example (id integer primary key autoincrement, value integer)",
477        )
478        .execute(connection)
479        .unwrap();
480
481        my_sum_utils::register_impl_with_behavior::<MySum, _>(
482            connection,
483            SqliteFunctionBehavior::DETERMINISTIC,
484        )
485        .unwrap();
486
487        let result = my_sum_example
488            .select(my_sum(value))
489            .get_result::<i32>(connection);
490        assert_eq!(Ok(0), result);
491    }
492
493    #[derive(Default)]
494    struct RangeMax<T> {
495        max_value: Option<T>,
496    }
497
498    impl<T: Default + Ord + Copy + Clone> SqliteAggregateFunction<(T, T, T)> for RangeMax<T> {
499        type Output = Option<T>;
500
501        fn step(&mut self, (x0, x1, x2): (T, T, T)) {
502            let max = if x0 >= x1 && x0 >= x2 {
503                x0
504            } else if x1 >= x0 && x1 >= x2 {
505                x1
506            } else {
507                x2
508            };
509
510            self.max_value = match self.max_value {
511                Some(current_max_value) if max > current_max_value => Some(max),
512                None => Some(max),
513                _ => self.max_value,
514            };
515        }
516
517        fn finalize(aggregator: Option<Self>) -> Self::Output {
518            aggregator?.max_value
519        }
520    }
521
522    table! {
523        range_max_example {
524            id -> Integer,
525            value1 -> Integer,
526            value2 -> Integer,
527            value3 -> Integer,
528        }
529    }
530
531    #[diesel_test_helper::test]
532    fn register_aggregate_multiarg_function() {
533        use self::range_max_example::dsl::*;
534
535        let connection = &mut connection();
536        crate::sql_query(
537            r#"CREATE TABLE range_max_example (
538                id integer primary key autoincrement,
539                value1 integer,
540                value2 integer,
541                value3 integer
542            )"#,
543        )
544        .execute(connection)
545        .unwrap();
546        crate::sql_query(
547            "INSERT INTO range_max_example (value1, value2, value3) VALUES (3, 2, 1), (2, 2, 2)",
548        )
549        .execute(connection)
550        .unwrap();
551
552        range_max_utils::register_impl_with_behavior::<RangeMax<i32>, _, _, _>(
553            connection,
554            SqliteFunctionBehavior::DETERMINISTIC,
555        )
556        .unwrap();
557        let result = range_max_example
558            .select(range_max(value1, value2, value3))
559            .get_result::<Option<i32>>(connection)
560            .unwrap();
561        assert_eq!(Some(3), result);
562    }
563
564    table! {
565        my_collation_example {
566            id -> Integer,
567            value -> Text,
568        }
569    }
570
571    #[diesel_test_helper::test]
572    fn register_collation_function() {
573        use self::my_collation_example::dsl::*;
574
575        let connection = &mut connection();
576
577        connection
578            .register_collation("RUSTNOCASE", |rhs, lhs| {
579                rhs.to_lowercase().cmp(&lhs.to_lowercase())
580            })
581            .unwrap();
582
583        crate::sql_query(
584                "CREATE TABLE my_collation_example (id integer primary key autoincrement, value text collate RUSTNOCASE)",
585            ).execute(connection)
586            .unwrap();
587        crate::sql_query(
588            "INSERT INTO my_collation_example (value) VALUES ('foo'), ('FOo'), ('f00')",
589        )
590        .execute(connection)
591        .unwrap();
592
593        let result = my_collation_example
594            .filter(value.eq("foo"))
595            .select(value)
596            .load::<String>(connection);
597        assert_eq!(
598            Ok(&["foo".to_owned(), "FOo".to_owned()][..]),
599            result.as_ref().map(|vec| vec.as_ref())
600        );
601
602        let result = my_collation_example
603            .filter(value.eq("FOO"))
604            .select(value)
605            .load::<String>(connection);
606        assert_eq!(
607            Ok(&["foo".to_owned(), "FOo".to_owned()][..]),
608            result.as_ref().map(|vec| vec.as_ref())
609        );
610
611        let result = my_collation_example
612            .filter(value.eq("f00"))
613            .select(value)
614            .load::<String>(connection);
615        assert_eq!(
616            Ok(&["f00".to_owned()][..]),
617            result.as_ref().map(|vec| vec.as_ref())
618        );
619
620        let result = my_collation_example
621            .filter(value.eq("F00"))
622            .select(value)
623            .load::<String>(connection);
624        assert_eq!(
625            Ok(&["f00".to_owned()][..]),
626            result.as_ref().map(|vec| vec.as_ref())
627        );
628
629        let result = my_collation_example
630            .filter(value.eq("oof"))
631            .select(value)
632            .load::<String>(connection);
633        assert_eq!(Ok(&[][..]), result.as_ref().map(|vec| vec.as_ref()));
634    }
635
636    #[diesel_test_helper::test]
637    fn aggregate_function_works_with_aligned_data() {
638        #[derive(Debug, Default)]
639        #[repr(align(64))]
640        struct OverAligned;
641
642        impl SqliteAggregateFunction<i32> for OverAligned {
643            type Output = i64;
644
645            fn step(&mut self, _value: i32) {
646                let need = core::mem::align_of::<Self>();
647                let got = core::mem::align_of_val(self);
648                assert_eq!(need, got);
649            }
650
651            fn finalize(_agg: Option<Self>) -> i64 {
652                0
653            }
654        }
655        #[declare_sql_function]
656        extern "SQL" {
657            #[aggregate]
658            fn over_aligned_sum(x: Integer) -> diesel::sql_types::BigInt;
659        }
660
661        let mut conn = SqliteConnection::establish(":memory:").unwrap();
662        over_aligned_sum_utils::register_impl::<OverAligned, _>(&mut conn).unwrap();
663
664        diesel::select(over_aligned_sum(1))
665            .execute(&mut conn)
666            .unwrap();
667    }
668
669    #[diesel_test_helper::test]
670    fn sum_twice() {
671        #[derive(Default)]
672        struct Sum(i32);
673
674        impl SqliteAggregateFunction<i32> for Sum {
675            type Output = i32;
676
677            fn step(&mut self, value: i32) {
678                self.0 += value;
679            }
680
681            fn finalize(agg: Option<Self>) -> i32 {
682                agg.map(|s| s.0).unwrap_or_default()
683            }
684        }
685
686        #[declare_sql_function]
687        extern "SQL" {
688            #[aggregate]
689            fn my_sum(x: Integer) -> Integer;
690        }
691
692        let mut conn = SqliteConnection::establish(":memory:").unwrap();
693        my_sum_utils::register_impl::<Sum, _>(&mut conn).unwrap();
694
695        conn.batch_execute(
696            "
697            CREATE TABLE test(key1 INTEGER, key2 INTEGER);
698            INSERT INTO test(key1, key2) VALUES (1, 2), (2, 4), (3, 6);
699",
700        )
701        .unwrap();
702
703        table! {
704            test (key1, key2) {
705                key1 -> Integer,
706                key2 -> Integer,
707            }
708        }
709
710        let (first_res, second_res) = test::table
711            .select((my_sum(test::key1), my_sum(test::key2)))
712            .get_result::<(i32, i32)>(&mut conn)
713            .unwrap();
714
715        assert_eq!(first_res, 6);
716        assert_eq!(second_res, 12);
717
718        conn.batch_execute("DELETE FROM test").unwrap();
719        let (first_res, second_res) = test::table
720            .select((my_sum(test::key1), my_sum(test::key2)))
721            .get_result::<(i32, i32)>(&mut conn)
722            .unwrap();
723
724        assert_eq!(first_res, 0);
725        assert_eq!(second_res, 0);
726    }
727
728    // ---- DIRECTONLY / INNOCUOUS function behavior tests ----
729
730    #[declare_sql_function]
731    extern "SQL" {
732        fn directonly_fn() -> Integer;
733        fn innocuous_fn() -> Integer;
734    }
735
736    #[diesel_test_helper::test]
737    fn directonly_function_blocked_from_view() {
738        let conn = &mut connection();
739
740        // Register a DIRECTONLY function
741        directonly_fn_utils::register_impl_with_behavior(
742            conn,
743            SqliteFunctionBehavior::DIRECTONLY,
744            || 42,
745        )
746        .unwrap();
747
748        // Direct call works
749        let result = crate::select(directonly_fn()).get_result::<i32>(conn);
750        assert_eq!(Ok(42), result);
751
752        // Create a view that calls the function
753        crate::sql_query("CREATE VIEW test_view AS SELECT directonly_fn() AS val")
754            .execute(conn)
755            .unwrap();
756
757        // Disable trusted schema so DIRECTONLY is enforced from schema objects
758        conn.set_trusted_schema(false).unwrap();
759
760        // Querying the view should fail because the function is DIRECTONLY
761        let result = crate::sql_query("SELECT val FROM test_view").execute(conn);
762        assert!(result.is_err());
763    }
764
765    #[diesel_test_helper::test]
766    fn innocuous_function_allowed_from_view_with_untrusted_schema() {
767        let conn = &mut connection();
768
769        // Register an INNOCUOUS function
770        innocuous_fn_utils::register_impl_with_behavior(
771            conn,
772            SqliteFunctionBehavior::DETERMINISTIC | SqliteFunctionBehavior::INNOCUOUS,
773            || 99,
774        )
775        .unwrap();
776
777        // Create a view that calls the function
778        crate::sql_query("CREATE VIEW innocuous_view AS SELECT innocuous_fn() AS val")
779            .execute(conn)
780            .unwrap();
781
782        // Disable trusted schema
783        conn.set_trusted_schema(false).unwrap();
784
785        // Querying the view should succeed because the function is INNOCUOUS
786        let result = crate::sql_query("SELECT val FROM innocuous_view").execute(conn);
787        assert!(result.is_ok());
788    }
789}