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#[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)] 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 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 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#[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 #[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 directonly_fn_utils::register_impl_with_behavior(
742 conn,
743 SqliteFunctionBehavior::DIRECTONLY,
744 || 42,
745 )
746 .unwrap();
747
748 let result = crate::select(directonly_fn()).get_result::<i32>(conn);
750 assert_eq!(Ok(42), result);
751
752 crate::sql_query("CREATE VIEW test_view AS SELECT directonly_fn() AS val")
754 .execute(conn)
755 .unwrap();
756
757 conn.set_trusted_schema(false).unwrap();
759
760 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 innocuous_fn_utils::register_impl_with_behavior(
771 conn,
772 SqliteFunctionBehavior::DETERMINISTIC | SqliteFunctionBehavior::INNOCUOUS,
773 || 99,
774 )
775 .unwrap();
776
777 crate::sql_query("CREATE VIEW innocuous_view AS SELECT innocuous_fn() AS val")
779 .execute(conn)
780 .unwrap();
781
782 conn.set_trusted_schema(false).unwrap();
784
785 let result = crate::sql_query("SELECT val FROM innocuous_view").execute(conn);
787 assert!(result.is_ok());
788 }
789}