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
7mod bind_collector;
8mod functions;
9#[cfg(all(test, not(all(target_family = "wasm", target_os = "unknown"))))]
10#[allow(unsafe_code)]
11mod oom_test_support;
12mod owned_row;
13mod raw;
14mod row;
15mod serialized_database;
16mod sqlite_value;
17mod statement_iterator;
18mod stmt;
19
20pub(in crate::sqlite) use self::bind_collector::SqliteBindCollector;
21pub use self::bind_collector::SqliteBindValue;
22pub use self::serialized_database::SerializedDatabase;
23pub use self::sqlite_value::SqliteValue;
24
25use std::os::raw as libc;
26
27use self::raw::RawConnection;
28use self::statement_iterator::*;
29use self::stmt::{Statement, StatementUse};
30use super::SqliteAggregateFunction;
31use crate::connection::instrumentation::{DynInstrumentation, StrQueryHelper};
32use crate::connection::statement_cache::StatementCache;
33use crate::connection::*;
34use crate::deserialize::{FromSqlRow, StaticallySizedRow};
35use crate::expression::QueryMetadata;
36use crate::query_builder::*;
37use crate::result::*;
38use crate::serialize::ToSql;
39use crate::sql_types::{HasSqlType, TypeMetadata};
40use crate::sqlite::Sqlite;
41
42#[allow(missing_debug_implementations)]
165#[cfg(feature = "sqlite")]
166pub struct SqliteConnection {
167 statement_cache: StatementCache<Sqlite, Statement>,
171 raw_connection: RawConnection,
172 transaction_state: AnsiTransactionManager,
173 metadata_lookup: (),
176 instrumentation: DynInstrumentation,
177 serialized_data: Vec<Vec<u8>>,
188}
189
190#[allow(unsafe_code)]
194unsafe impl Send for SqliteConnection {}
195
196impl SimpleConnection for SqliteConnection {
197 fn batch_execute(&mut self, query: &str) -> QueryResult<()> {
198 self.instrumentation
199 .on_connection_event(InstrumentationEvent::StartQuery {
200 query: &StrQueryHelper::new(query),
201 });
202 let resp = self.raw_connection.exec(query);
203 self.instrumentation
204 .on_connection_event(InstrumentationEvent::FinishQuery {
205 query: &StrQueryHelper::new(query),
206 error: resp.as_ref().err(),
207 });
208 if resp.is_err() && self.raw_connection.is_autocommit() {
209 self.transaction_state.status = TransactionManagerStatus::Valid(Default::default());
211 }
212 resp
213 }
214}
215
216impl ConnectionSealed for SqliteConnection {}
217
218impl Connection for SqliteConnection {
219 type Backend = Sqlite;
220 type TransactionManager = AnsiTransactionManager;
221
222 fn establish(database_url: &str) -> ConnectionResult<Self> {
240 let mut instrumentation = DynInstrumentation::default_instrumentation();
241 instrumentation.on_connection_event(InstrumentationEvent::StartEstablishConnection {
242 url: database_url,
243 });
244
245 let establish_result = Self::establish_inner(database_url);
246 instrumentation.on_connection_event(InstrumentationEvent::FinishEstablishConnection {
247 url: database_url,
248 error: establish_result.as_ref().err(),
249 });
250 let mut conn = establish_result?;
251 conn.instrumentation = instrumentation;
252 Ok(conn)
253 }
254
255 fn execute_returning_count<T>(&mut self, source: &T) -> QueryResult<usize>
256 where
257 T: QueryFragment<Self::Backend> + QueryId,
258 {
259 let statement_use = self.prepared_query(source)?;
260 statement_use.run().and_then(|_| {
261 self.raw_connection
262 .rows_affected_by_last_query()
263 .map_err(Error::DeserializationError)
264 })
265 }
266
267 fn transaction_state(&mut self) -> &mut AnsiTransactionManager
268 where
269 Self: Sized,
270 {
271 &mut self.transaction_state
272 }
273
274 fn instrumentation(&mut self) -> &mut dyn Instrumentation {
275 &mut *self.instrumentation
276 }
277
278 fn set_instrumentation(&mut self, instrumentation: impl Instrumentation) {
279 self.instrumentation = instrumentation.into();
280 }
281
282 fn set_prepared_statement_cache_size(&mut self, size: CacheSize) {
283 self.statement_cache.set_cache_size(size);
284 }
285}
286
287impl LoadConnection<DefaultLoadingMode> for SqliteConnection {
288 type Cursor<'conn, 'query> = StatementIterator<'conn, 'query>;
289 type Row<'conn, 'query> = self::row::SqliteRow<'conn, 'query>;
290
291 fn load<'conn, 'query, T>(
292 &'conn mut self,
293 source: T,
294 ) -> QueryResult<Self::Cursor<'conn, 'query>>
295 where
296 T: Query + QueryFragment<Self::Backend> + QueryId + 'query,
297 Self::Backend: QueryMetadata<T::SqlType>,
298 {
299 let statement = self.prepared_query(source)?;
300
301 Ok(StatementIterator::new(statement))
302 }
303}
304
305impl WithMetadataLookup for SqliteConnection {
306 fn metadata_lookup(&mut self) -> &mut <Sqlite as TypeMetadata>::MetadataLookup {
307 &mut self.metadata_lookup
308 }
309}
310
311#[cfg(feature = "r2d2")]
312impl crate::r2d2::R2D2Connection for crate::sqlite::SqliteConnection {
313 fn ping(&mut self) -> QueryResult<()> {
314 use crate::RunQueryDsl;
315
316 crate::r2d2::CheckConnectionQuery.execute(self).map(|_| ())
317 }
318
319 fn is_broken(&mut self) -> bool {
320 AnsiTransactionManager::is_broken_transaction_manager(self)
321 }
322}
323
324impl MultiConnectionHelper for SqliteConnection {
325 fn to_any<'a>(
326 lookup: &mut <Self::Backend as crate::sql_types::TypeMetadata>::MetadataLookup,
327 ) -> &mut (dyn std::any::Any + 'a) {
328 lookup
329 }
330
331 fn from_any(
332 lookup: &mut dyn std::any::Any,
333 ) -> Option<&mut <Self::Backend as crate::sql_types::TypeMetadata>::MetadataLookup> {
334 lookup.downcast_mut()
335 }
336}
337
338impl SqliteConnection {
339 pub fn immediate_transaction<T, E, F>(&mut self, f: F) -> Result<T, E>
361 where
362 F: FnOnce(&mut Self) -> Result<T, E>,
363 E: From<Error>,
364 {
365 self.transaction_sql(f, "BEGIN IMMEDIATE")
366 }
367
368 pub fn exclusive_transaction<T, E, F>(&mut self, f: F) -> Result<T, E>
390 where
391 F: FnOnce(&mut Self) -> Result<T, E>,
392 E: From<Error>,
393 {
394 self.transaction_sql(f, "BEGIN EXCLUSIVE")
395 }
396
397 fn transaction_sql<T, E, F>(&mut self, f: F, sql: &str) -> Result<T, E>
398 where
399 F: FnOnce(&mut Self) -> Result<T, E>,
400 E: From<Error>,
401 {
402 AnsiTransactionManager::begin_transaction_sql(&mut *self, sql)?;
403 match f(&mut *self) {
404 Ok(value) => {
405 AnsiTransactionManager::commit_transaction(&mut *self)?;
406 Ok(value)
407 }
408 Err(e) => {
409 AnsiTransactionManager::rollback_transaction(&mut *self)?;
410 Err(e)
411 }
412 }
413 }
414
415 fn prepared_query<'conn, 'query, T>(
416 &'conn mut self,
417 source: T,
418 ) -> QueryResult<StatementUse<'conn, 'query>>
419 where
420 T: QueryFragment<Sqlite> + QueryId + 'query,
421 {
422 self.instrumentation
423 .on_connection_event(InstrumentationEvent::StartQuery {
424 query: &crate::debug_query(&source),
425 });
426 let raw_connection = &self.raw_connection;
427 let cache = &mut self.statement_cache;
428 let statement = match cache.cached_statement(
429 &source,
430 &Sqlite,
431 &[],
432 raw_connection,
433 Statement::prepare,
434 &mut *self.instrumentation,
435 ) {
436 Ok(statement) => statement,
437 Err(e) => {
438 self.instrumentation
439 .on_connection_event(InstrumentationEvent::FinishQuery {
440 query: &crate::debug_query(&source),
441 error: Some(&e),
442 });
443
444 return Err(e);
445 }
446 };
447
448 StatementUse::bind(statement, source, &mut *self.instrumentation)
449 }
450
451 #[doc(hidden)]
452 pub fn register_sql_function<ArgsSqlType, RetSqlType, Args, Ret, F>(
453 &mut self,
454 fn_name: &str,
455 deterministic: bool,
456 mut f: F,
457 ) -> QueryResult<()>
458 where
459 F: FnMut(Args) -> Ret + std::panic::UnwindSafe + Send + 'static,
460 Args: FromSqlRow<ArgsSqlType, Sqlite> + StaticallySizedRow<ArgsSqlType, Sqlite>,
461 Ret: ToSql<RetSqlType, Sqlite>,
462 Sqlite: HasSqlType<RetSqlType>,
463 {
464 functions::register(
465 &self.raw_connection,
466 fn_name,
467 deterministic,
468 move |_, args| f(args),
469 )
470 }
471
472 #[doc(hidden)]
473 pub fn register_noarg_sql_function<RetSqlType, Ret, F>(
474 &self,
475 fn_name: &str,
476 deterministic: bool,
477 f: F,
478 ) -> QueryResult<()>
479 where
480 F: FnMut() -> Ret + std::panic::UnwindSafe + Send + 'static,
481 Ret: ToSql<RetSqlType, Sqlite>,
482 Sqlite: HasSqlType<RetSqlType>,
483 {
484 functions::register_noargs(&self.raw_connection, fn_name, deterministic, f)
485 }
486
487 #[doc(hidden)]
488 pub fn register_aggregate_function<ArgsSqlType, RetSqlType, Args, Ret, A>(
489 &mut self,
490 fn_name: &str,
491 ) -> QueryResult<()>
492 where
493 A: SqliteAggregateFunction<Args, Output = Ret> + 'static + Send + std::panic::UnwindSafe,
494 Args: FromSqlRow<ArgsSqlType, Sqlite> + StaticallySizedRow<ArgsSqlType, Sqlite>,
495 Ret: ToSql<RetSqlType, Sqlite>,
496 Sqlite: HasSqlType<RetSqlType>,
497 {
498 functions::register_aggregate::<_, _, _, _, A>(&self.raw_connection, fn_name)
499 }
500
501 pub fn register_collation<F>(&mut self, collation_name: &str, collation: F) -> QueryResult<()>
537 where
538 F: Fn(&str, &str) -> std::cmp::Ordering + Send + 'static + std::panic::UnwindSafe,
539 {
540 self.raw_connection
541 .register_collation_function(collation_name, collation)
542 }
543
544 pub fn serialize_database_to_buffer(&mut self) -> SerializedDatabase {
555 self.raw_connection.serialize()
556 }
557
558 #[allow(unsafe_code)]
598 pub fn deserialize_readonly_database_from_buffer(&mut self, data: &[u8]) -> QueryResult<()> {
599 self.serialized_data.push(data.to_vec());
602 let last = self
603 .serialized_data
604 .last()
605 .expect("We literally pushed it above, so it's there");
606 unsafe {
607 self.raw_connection.deserialize(last)
610 }
611 }
612
613 fn register_diesel_sql_functions(&self) -> QueryResult<()> {
614 use crate::sql_types::{Integer, Text};
615
616 functions::register::<Text, Integer, _, _, _>(
620 &self.raw_connection,
621 "diesel_manage_updated_at",
622 false,
623 |conn, table_name: String| {
624 conn.exec(&::alloc::__export::must_use({
::alloc::fmt::format(format_args!("CREATE TRIGGER __diesel_manage_updated_at_{0}\nAFTER UPDATE ON {0}\nFOR EACH ROW WHEN\n old.updated_at IS NULL AND\n new.updated_at IS NULL OR\n old.updated_at == new.updated_at\nBEGIN\n UPDATE {0}\n SET updated_at = CURRENT_TIMESTAMP\n WHERE ROWID = new.ROWID;\nEND\n",
table_name))
})format!(
625 include_str!("diesel_manage_updated_at.sql"),
626 table_name = table_name
627 ))
628 .expect("Failed to create trigger");
629 0 },
631 )
632 }
633
634 fn establish_inner(database_url: &str) -> Result<SqliteConnection, ConnectionError> {
635 use crate::result::ConnectionError::CouldntSetupConfiguration;
636 let raw_connection = RawConnection::establish(database_url)?;
637 let conn = Self {
638 statement_cache: StatementCache::new(),
639 raw_connection,
640 transaction_state: AnsiTransactionManager::default(),
641 metadata_lookup: (),
642 instrumentation: DynInstrumentation::none(),
643 serialized_data: Vec::new(),
644 };
645 conn.register_diesel_sql_functions()
646 .map_err(CouldntSetupConfiguration)?;
647 Ok(conn)
648 }
649}
650
651fn error_message(err_code: libc::c_int) -> &'static str {
652 ffi::code_to_str(err_code)
653}
654
655#[cfg(test)]
656mod tests {
657 use super::*;
658 use crate::dsl::sql;
659 use crate::prelude::*;
660 use crate::sql_types::{Integer, Text};
661 use crate::test_helpers::format_error;
662
663 fn connection() -> SqliteConnection {
664 SqliteConnection::establish(":memory:").unwrap()
665 }
666
667 #[declare_sql_function]
668 extern "SQL" {
669 fn fun_case(x: Text) -> Text;
670 fn my_add(x: Integer, y: Integer) -> Integer;
671 fn answer() -> Integer;
672 fn add_counter(x: Integer) -> Integer;
673
674 #[aggregate]
675 fn my_sum(expr: Integer) -> Integer;
676 #[aggregate]
677 fn range_max(expr1: Integer, expr2: Integer, expr3: Integer) -> Nullable<Integer>;
678 }
679
680 #[diesel_test_helper::test]
681 fn database_serializes_and_deserializes_successfully() {
682 let expected_users = vec![
683 (
684 1,
685 "John Doe".to_string(),
686 "john.doe@example.com".to_string(),
687 ),
688 (
689 2,
690 "Jane Doe".to_string(),
691 "jane.doe@example.com".to_string(),
692 ),
693 ];
694
695 let conn1 = &mut connection();
696 let _ =
697 crate::sql_query("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT, email TEXT)")
698 .execute(conn1);
699 let _ = crate::sql_query("INSERT INTO users (name, email) VALUES ('John Doe', 'john.doe@example.com'), ('Jane Doe', 'jane.doe@example.com')")
700 .execute(conn1);
701
702 for _i in 0..2 {
703 let serialized_database = conn1.serialize_database_to_buffer();
704 let conn2 = &mut connection();
705 conn2
706 .deserialize_readonly_database_from_buffer(
707 serialized_database.try_as_slice().unwrap(),
708 )
709 .unwrap();
710
711 let query =
712 sql::<(Integer, Text, Text)>("SELECT id, name, email FROM users ORDER BY id");
713 let actual_users = query.load::<(i32, String, String)>(conn2).unwrap();
714
715 assert_eq!(expected_users, actual_users);
716 std::mem::drop(serialized_database);
720 let query =
721 sql::<(Integer, Text, Text)>("SELECT id, name, email FROM users ORDER BY id");
722 let actual_users = query.load::<(i32, String, String)>(conn2).unwrap();
723
724 assert_eq!(expected_users, actual_users);
725 }
726 }
727
728 #[diesel_test_helper::test]
729 fn database_deserialize_random_bytes() {
730 let buffer = vec![0, 1, 2, 3, 4];
731 let conn = &mut SqliteConnection::establish(":memory:").unwrap();
732
733 conn.deserialize_readonly_database_from_buffer(&buffer)
734 .unwrap();
735
736 let r = sql::<Integer>("SELECT id FROM users").load::<i32>(conn);
737
738 assert!(r.is_err());
739 assert_eq!(format_error(&r.unwrap_err()), "file is not a database");
740
741 let conn = &mut SqliteConnection::establish(":memory:").unwrap();
742
743 let _ =
744 crate::sql_query("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT, email TEXT)")
745 .execute(conn);
746 let _ = crate::sql_query("INSERT INTO users (name, email) VALUES ('John Doe', 'john.doe@example.com'), ('Jane Doe', 'jane.doe@example.com')")
747 .execute(conn);
748
749 let db = conn.serialize_database_to_buffer();
750 let mut bad_buffer = db[..100].to_vec();
752 bad_buffer.extend(b"whatever");
753 conn.deserialize_readonly_database_from_buffer(&bad_buffer)
754 .unwrap();
755
756 let r = sql::<Integer>("SELECT id FROM users").load::<i32>(conn);
757
758 assert!(r.is_err());
759 assert_eq!(
760 format_error(&r.unwrap_err()),
761 "database disk image is malformed"
762 );
763
764 let mut size_fitting_bad_buffer = db[..100].to_vec();
766 size_fitting_bad_buffer.extend(
767 core::iter::repeat(b"abcdefghij")
768 .flatten()
769 .take(db.len() - 100),
770 );
771 let r = conn.deserialize_readonly_database_from_buffer(&size_fitting_bad_buffer);
772
773 assert!(r.is_err());
774 assert_eq!(
775 format_error(&r.unwrap_err()),
776 "database disk image is malformed"
777 );
778 }
779
780 #[diesel_test_helper::test]
781 fn database_serializes_empty_deserialized_database() {
782 let conn = &mut SqliteConnection::establish(":memory:").unwrap();
783 conn.deserialize_readonly_database_from_buffer(&[]).unwrap();
784
785 let serialized = conn.serialize_database_to_buffer();
786
787 assert!(serialized.is_empty());
788 assert!(serialized.try_as_slice().unwrap().is_empty());
789 }
790
791 #[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
792 #[allow(unsafe_code)]
793 mod sqlite_serialize_oom {
794 use super::super::oom_test_support::{panic_message, run_in_child, with_heap_limit};
795 use super::super::{ffi, SerializedDatabase};
796 use crate::connection::{Connection, SimpleConnection};
797 use crate::sqlite::SqliteConnection;
798 use crate::test_helpers::format_error;
799
800 const MIN_DATABASE_BYTES: i64 = 1_048_576;
801
802 fn with_failing_serialize<R>(f: impl FnOnce() -> R) -> R {
805 with_heap_limit(65_536, f)
806 }
807
808 #[test]
809 fn sqlite_serialize_oom_is_contained() {
810 run_in_child(|| {
811 let mut conn = large_database();
812
813 let (baseline_size, baseline) = serialize_direct(&conn);
814 assert!(
815 baseline_size >= MIN_DATABASE_BYTES,
816 "the serialized database is smaller than 1 MiB"
817 );
818 assert!(
819 !baseline.is_null(),
820 "SQLite refused to serialize a valid database"
821 );
822 unsafe { ffi::sqlite3_free(baseline as _) };
824
825 let (reported_size, data) = with_failing_serialize(|| serialize_direct(&conn));
826 if !data.is_null() {
827 unsafe { ffi::sqlite3_free(data as _) };
829 }
830 assert!(
831 data.is_null(),
832 "SQLite did not fail the output allocation of the serialization"
833 );
834 assert!(
836 reported_size >= MIN_DATABASE_BYTES,
837 "SQLite reported a serialization size of {reported_size} with a null buffer"
838 );
839
840 let serialized: SerializedDatabase =
841 with_failing_serialize(|| conn.serialize_database_to_buffer());
842 let error = serialized
843 .try_as_slice()
844 .expect_err("the failed output allocation must surface as an error");
845 assert_eq!(format_error(&error), "out of memory");
846
847 let payload = std::panic::catch_unwind(core::panic::AssertUnwindSafe(|| {
848 core::hint::black_box(serialized[0]);
849 }))
850 .expect_err("the serialized database access did not panic");
851 let message = panic_message(&*payload);
852 assert!(
853 message.contains("Cannot access the serialized database: out of memory"),
854 "SQLite serialization allocation failure surfaced as `{message}` instead \
855 of a caught allocation panic"
856 );
857 });
858 }
859
860 fn large_database() -> SqliteConnection {
861 let mut conn = SqliteConnection::establish(":memory:").unwrap();
862 conn.batch_execute(&format!(
863 "CREATE TABLE blobs (id INTEGER PRIMARY KEY, payload BLOB);
864 INSERT INTO blobs (payload) VALUES (zeroblob({MIN_DATABASE_BYTES}));"
865 ))
866 .unwrap();
867 conn
868 }
869
870 fn serialize_direct(conn: &SqliteConnection) -> (ffi::sqlite3_int64, *mut u8) {
871 unsafe {
873 let mut size: ffi::sqlite3_int64 = 0;
874 let data = ffi::sqlite3_serialize(
875 conn.raw_connection.internal_connection.as_ptr(),
876 core::ptr::null(),
877 &mut size as *mut _,
878 0,
879 );
880 (size, data)
881 }
882 }
883 }
884
885 #[diesel_test_helper::test]
886 fn register_custom_function() {
887 let connection = &mut connection();
888 fun_case_utils::register_impl(connection, |x: String| {
889 x.chars()
890 .enumerate()
891 .map(|(i, c)| {
892 if i % 2 == 0 {
893 c.to_lowercase().to_string()
894 } else {
895 c.to_uppercase().to_string()
896 }
897 })
898 .collect::<String>()
899 })
900 .unwrap();
901
902 let mapped_string = crate::select(fun_case("foobar"))
903 .get_result::<String>(connection)
904 .unwrap();
905 assert_eq!("fOoBaR", mapped_string);
906 }
907
908 #[diesel_test_helper::test]
909 fn register_multiarg_function() {
910 let connection = &mut connection();
911 my_add_utils::register_impl(connection, |x: i32, y: i32| x + y).unwrap();
912
913 let added = crate::select(my_add(1, 2)).get_result::<i32>(connection);
914 assert_eq!(Ok(3), added);
915 }
916
917 #[diesel_test_helper::test]
918 fn register_noarg_function() {
919 let connection = &mut connection();
920 answer_utils::register_impl(connection, || 42).unwrap();
921
922 let answer = crate::select(answer()).get_result::<i32>(connection);
923 assert_eq!(Ok(42), answer);
924 }
925
926 #[diesel_test_helper::test]
927 fn register_nondeterministic_noarg_function() {
928 let connection = &mut connection();
929 answer_utils::register_nondeterministic_impl(connection, || 42).unwrap();
930
931 let answer = crate::select(answer()).get_result::<i32>(connection);
932 assert_eq!(Ok(42), answer);
933 }
934
935 #[diesel_test_helper::test]
936 fn register_nondeterministic_function() {
937 let connection = &mut connection();
938 let mut y = 0;
939 add_counter_utils::register_nondeterministic_impl(connection, move |x: i32| {
940 y += 1;
941 x + y
942 })
943 .unwrap();
944
945 let added = crate::select((add_counter(1), add_counter(1), add_counter(1)))
946 .get_result::<(i32, i32, i32)>(connection);
947 assert_eq!(Ok((2, 3, 4)), added);
948 }
949
950 #[derive(Default)]
951 struct MySum {
952 sum: i32,
953 }
954
955 impl SqliteAggregateFunction<i32> for MySum {
956 type Output = i32;
957
958 fn step(&mut self, expr: i32) {
959 self.sum += expr;
960 }
961
962 fn finalize(aggregator: Option<Self>) -> Self::Output {
963 aggregator.map(|a| a.sum).unwrap_or_default()
964 }
965 }
966
967 table! {
968 my_sum_example {
969 id -> Integer,
970 value -> Integer,
971 }
972 }
973
974 #[diesel_test_helper::test]
975 fn register_aggregate_function() {
976 use self::my_sum_example::dsl::*;
977
978 let connection = &mut connection();
979 crate::sql_query(
980 "CREATE TABLE my_sum_example (id integer primary key autoincrement, value integer)",
981 )
982 .execute(connection)
983 .unwrap();
984 crate::sql_query("INSERT INTO my_sum_example (value) VALUES (1), (2), (3)")
985 .execute(connection)
986 .unwrap();
987
988 my_sum_utils::register_impl::<MySum, _>(connection).unwrap();
989
990 let result = my_sum_example
991 .select(my_sum(value))
992 .get_result::<i32>(connection);
993 assert_eq!(Ok(6), result);
994 }
995
996 #[diesel_test_helper::test]
997 fn register_aggregate_function_returns_finalize_default_on_empty_set() {
998 use self::my_sum_example::dsl::*;
999
1000 let connection = &mut connection();
1001 crate::sql_query(
1002 "CREATE TABLE my_sum_example (id integer primary key autoincrement, value integer)",
1003 )
1004 .execute(connection)
1005 .unwrap();
1006
1007 my_sum_utils::register_impl::<MySum, _>(connection).unwrap();
1008
1009 let result = my_sum_example
1010 .select(my_sum(value))
1011 .get_result::<i32>(connection);
1012 assert_eq!(Ok(0), result);
1013 }
1014
1015 #[derive(Default)]
1016 struct RangeMax<T> {
1017 max_value: Option<T>,
1018 }
1019
1020 impl<T: Default + Ord + Copy + Clone> SqliteAggregateFunction<(T, T, T)> for RangeMax<T> {
1021 type Output = Option<T>;
1022
1023 fn step(&mut self, (x0, x1, x2): (T, T, T)) {
1024 let max = if x0 >= x1 && x0 >= x2 {
1025 x0
1026 } else if x1 >= x0 && x1 >= x2 {
1027 x1
1028 } else {
1029 x2
1030 };
1031
1032 self.max_value = match self.max_value {
1033 Some(current_max_value) if max > current_max_value => Some(max),
1034 None => Some(max),
1035 _ => self.max_value,
1036 };
1037 }
1038
1039 fn finalize(aggregator: Option<Self>) -> Self::Output {
1040 aggregator?.max_value
1041 }
1042 }
1043
1044 table! {
1045 range_max_example {
1046 id -> Integer,
1047 value1 -> Integer,
1048 value2 -> Integer,
1049 value3 -> Integer,
1050 }
1051 }
1052
1053 #[diesel_test_helper::test]
1054 fn register_aggregate_multiarg_function() {
1055 use self::range_max_example::dsl::*;
1056
1057 let connection = &mut connection();
1058 crate::sql_query(
1059 r#"CREATE TABLE range_max_example (
1060 id integer primary key autoincrement,
1061 value1 integer,
1062 value2 integer,
1063 value3 integer
1064 )"#,
1065 )
1066 .execute(connection)
1067 .unwrap();
1068 crate::sql_query(
1069 "INSERT INTO range_max_example (value1, value2, value3) VALUES (3, 2, 1), (2, 2, 2)",
1070 )
1071 .execute(connection)
1072 .unwrap();
1073
1074 range_max_utils::register_impl::<RangeMax<i32>, _, _, _>(connection).unwrap();
1075 let result = range_max_example
1076 .select(range_max(value1, value2, value3))
1077 .get_result::<Option<i32>>(connection)
1078 .unwrap();
1079 assert_eq!(Some(3), result);
1080 }
1081
1082 table! {
1083 my_collation_example {
1084 id -> Integer,
1085 value -> Text,
1086 }
1087 }
1088
1089 #[diesel_test_helper::test]
1090 fn register_collation_function() {
1091 use self::my_collation_example::dsl::*;
1092
1093 let connection = &mut connection();
1094
1095 connection
1096 .register_collation("RUSTNOCASE", |rhs, lhs| {
1097 rhs.to_lowercase().cmp(&lhs.to_lowercase())
1098 })
1099 .unwrap();
1100
1101 crate::sql_query(
1102 "CREATE TABLE my_collation_example (id integer primary key autoincrement, value text collate RUSTNOCASE)",
1103 ).execute(connection)
1104 .unwrap();
1105 crate::sql_query(
1106 "INSERT INTO my_collation_example (value) VALUES ('foo'), ('FOo'), ('f00')",
1107 )
1108 .execute(connection)
1109 .unwrap();
1110
1111 let result = my_collation_example
1112 .filter(value.eq("foo"))
1113 .select(value)
1114 .load::<String>(connection);
1115 assert_eq!(
1116 Ok(&["foo".to_owned(), "FOo".to_owned()][..]),
1117 result.as_ref().map(|vec| vec.as_ref())
1118 );
1119
1120 let result = my_collation_example
1121 .filter(value.eq("FOO"))
1122 .select(value)
1123 .load::<String>(connection);
1124 assert_eq!(
1125 Ok(&["foo".to_owned(), "FOo".to_owned()][..]),
1126 result.as_ref().map(|vec| vec.as_ref())
1127 );
1128
1129 let result = my_collation_example
1130 .filter(value.eq("f00"))
1131 .select(value)
1132 .load::<String>(connection);
1133 assert_eq!(
1134 Ok(&["f00".to_owned()][..]),
1135 result.as_ref().map(|vec| vec.as_ref())
1136 );
1137
1138 let result = my_collation_example
1139 .filter(value.eq("F00"))
1140 .select(value)
1141 .load::<String>(connection);
1142 assert_eq!(
1143 Ok(&["f00".to_owned()][..]),
1144 result.as_ref().map(|vec| vec.as_ref())
1145 );
1146
1147 let result = my_collation_example
1148 .filter(value.eq("oof"))
1149 .select(value)
1150 .load::<String>(connection);
1151 assert_eq!(Ok(&[][..]), result.as_ref().map(|vec| vec.as_ref()));
1152 }
1153
1154 #[diesel_test_helper::test]
1156 fn test_correct_serialization_of_owned_strings() {
1157 use crate::prelude::*;
1158
1159 #[derive(Debug, crate::expression::AsExpression)]
1160 #[diesel(sql_type = diesel::sql_types::Text)]
1161 struct CustomWrapper(String);
1162
1163 impl crate::serialize::ToSql<Text, Sqlite> for CustomWrapper {
1164 fn to_sql<'b>(
1165 &'b self,
1166 out: &mut crate::serialize::Output<'b, '_, Sqlite>,
1167 ) -> crate::serialize::Result {
1168 out.set_value(self.0.to_string());
1169 Ok(crate::serialize::IsNull::No)
1170 }
1171 }
1172
1173 let connection = &mut connection();
1174
1175 let res = crate::select(
1176 CustomWrapper("".into())
1177 .into_sql::<crate::sql_types::Text>()
1178 .nullable(),
1179 )
1180 .get_result::<Option<String>>(connection)
1181 .unwrap();
1182 assert_eq!(res, Some(String::new()));
1183 }
1184
1185 #[diesel_test_helper::test]
1186 fn test_correct_serialization_of_owned_bytes() {
1187 use crate::prelude::*;
1188
1189 #[derive(Debug, crate::expression::AsExpression)]
1190 #[diesel(sql_type = diesel::sql_types::Binary)]
1191 struct CustomWrapper(Vec<u8>);
1192
1193 impl crate::serialize::ToSql<crate::sql_types::Binary, Sqlite> for CustomWrapper {
1194 fn to_sql<'b>(
1195 &'b self,
1196 out: &mut crate::serialize::Output<'b, '_, Sqlite>,
1197 ) -> crate::serialize::Result {
1198 out.set_value(self.0.clone());
1199 Ok(crate::serialize::IsNull::No)
1200 }
1201 }
1202
1203 let connection = &mut connection();
1204
1205 let res = crate::select(
1206 CustomWrapper(Vec::new())
1207 .into_sql::<crate::sql_types::Binary>()
1208 .nullable(),
1209 )
1210 .get_result::<Option<Vec<u8>>>(connection)
1211 .unwrap();
1212 assert_eq!(res, Some(Vec::new()));
1213 }
1214
1215 #[diesel_test_helper::test]
1216 fn correctly_handle_empty_query() {
1217 let check_empty_query_error = |r: crate::QueryResult<usize>| {
1218 assert!(r.is_err());
1219 let err = r.unwrap_err();
1220 assert!(
1221 matches!(err, crate::result::Error::QueryBuilderError(ref b) if b.is::<crate::result::EmptyQuery>()),
1222 "Expected a query builder error, but got {err}"
1223 );
1224 };
1225 let connection = &mut SqliteConnection::establish(":memory:").unwrap();
1226 check_empty_query_error(crate::sql_query("").execute(connection));
1227 check_empty_query_error(crate::sql_query(" ").execute(connection));
1228 check_empty_query_error(crate::sql_query("\n\t").execute(connection));
1229 check_empty_query_error(crate::sql_query("-- SELECT 1;").execute(connection));
1230 }
1231
1232 #[diesel_test_helper::test]
1233 fn aggregate_function_works_with_aligned_data() {
1234 #[derive(Debug, Default)]
1235 #[repr(align(64))]
1236 struct OverAligned;
1237
1238 impl SqliteAggregateFunction<i32> for OverAligned {
1239 type Output = i64;
1240
1241 fn step(&mut self, _value: i32) {
1242 let need = core::mem::align_of::<Self>();
1243 let got = core::mem::align_of_val(self);
1244 assert_eq!(need, got);
1245 }
1246
1247 fn finalize(_agg: Option<Self>) -> i64 {
1248 0
1249 }
1250 }
1251 #[declare_sql_function]
1252 extern "SQL" {
1253 #[aggregate]
1254 fn over_aligned_sum(x: Integer) -> diesel::sql_types::BigInt;
1255 }
1256
1257 let mut conn = SqliteConnection::establish(":memory:").unwrap();
1258 over_aligned_sum_utils::register_impl::<OverAligned, _>(&mut conn).unwrap();
1259
1260 diesel::select(over_aligned_sum(1))
1261 .execute(&mut conn)
1262 .unwrap();
1263 }
1264
1265 #[diesel_test_helper::test]
1266 fn sum_twice() {
1267 #[derive(Default)]
1268 struct Sum(i32);
1269
1270 impl SqliteAggregateFunction<i32> for Sum {
1271 type Output = i32;
1272
1273 fn step(&mut self, value: i32) {
1274 self.0 += value;
1275 }
1276
1277 fn finalize(agg: Option<Self>) -> i32 {
1278 agg.map(|s| s.0).unwrap_or_default()
1279 }
1280 }
1281
1282 #[declare_sql_function]
1283 extern "SQL" {
1284 #[aggregate]
1285 fn my_sum(x: Integer) -> Integer;
1286 }
1287
1288 let mut conn = SqliteConnection::establish(":memory:").unwrap();
1289 my_sum_utils::register_impl::<Sum, _>(&mut conn).unwrap();
1290
1291 conn.batch_execute(
1292 "
1293 CREATE TABLE test(key1 INTEGER, key2 INTEGER);
1294 INSERT INTO test(key1, key2) VALUES (1, 2), (2, 4), (3, 6);
1295",
1296 )
1297 .unwrap();
1298
1299 table! {
1300 test (key1, key2) {
1301 key1 -> Integer,
1302 key2 -> Integer,
1303 }
1304 }
1305
1306 let (first_res, second_res) = test::table
1307 .select((my_sum(test::key1), my_sum(test::key2)))
1308 .get_result::<(i32, i32)>(&mut conn)
1309 .unwrap();
1310
1311 assert_eq!(first_res, 6);
1312 assert_eq!(second_res, 12);
1313
1314 conn.batch_execute("DELETE FROM test").unwrap();
1315 let (first_res, second_res) = test::table
1316 .select((my_sum(test::key1), my_sum(test::key2)))
1317 .get_result::<(i32, i32)>(&mut conn)
1318 .unwrap();
1319
1320 assert_eq!(first_res, 0);
1321 assert_eq!(second_res, 0);
1322 }
1323
1324 #[diesel_test_helper::test]
1325 fn test_injection() {
1326 diesel::table! {
1327 #[sql_name = "quote'table"]
1328 quote_table (id) {
1329 id -> Nullable<Integer>,
1330 name -> Nullable<Text>,
1331 }
1332 }
1333
1334 let mut conn = SqliteConnection::establish(":memory:").unwrap();
1335
1336 conn.batch_execute("CREATE TABLE \"quote'table\" (id INTEGER PRIMARY KEY, name TEXT);")
1337 .unwrap();
1338
1339 diesel::insert_into(quote_table::table)
1340 .values((quote_table::id.eq(1), quote_table::name.eq("Jane")))
1341 .execute(&mut conn)
1342 .unwrap();
1343
1344 let data = quote_table::table
1345 .load::<(Option<i32>, Option<String>)>(&mut conn)
1346 .unwrap();
1347 assert_eq!(data, [(Some(1), Some("Jane".to_owned()))]);
1348 }
1349}