Skip to main content

diesel/sqlite/connection/
row.rs

1use super::owned_row::OwnedSqliteRow;
2use super::sqlite_value::{OwnedSqliteValue, SqliteValue};
3use super::stmt::StatementUse;
4use crate::backend::Backend;
5use crate::result::QueryResult;
6use crate::row::{Field, IntoOwnedRow, PartialRow, Row, RowIndex, RowSealed};
7use crate::sqlite::Sqlite;
8use alloc::borrow::ToOwned;
9use alloc::rc::Rc;
10use alloc::string::String;
11use alloc::sync::Arc;
12use alloc::vec::Vec;
13use core::cell::{Ref, RefCell};
14
15#[allow(missing_debug_implementations)]
16pub struct SqliteRow<'stmt, 'query> {
17    pub(super) inner: Rc<RefCell<PrivateSqliteRow<'stmt, 'query>>>,
18    pub(super) field_count: usize,
19}
20
21pub(super) enum PrivateSqliteRow<'stmt, 'query> {
22    Direct(StatementUse<'stmt, 'query>),
23    Duplicated {
24        values: Vec<Option<OwnedSqliteValue>>,
25        column_names: Rc<[Option<String>]>,
26    },
27}
28
29impl<'stmt> IntoOwnedRow<'stmt, Sqlite> for SqliteRow<'stmt, '_> {
30    type OwnedRow = OwnedSqliteRow;
31
32    type Cache = Option<Arc<[Option<String>]>>;
33
34    fn into_owned(self, column_name_cache: &mut Self::Cache) -> Self::OwnedRow {
35        self.inner.borrow().moveable(column_name_cache)
36    }
37}
38
39impl<'stmt, 'query> PrivateSqliteRow<'stmt, 'query> {
40    pub(super) fn duplicate(
41        &mut self,
42        column_names: &mut Option<Rc<[Option<String>]>>,
43    ) -> QueryResult<PrivateSqliteRow<'stmt, 'query>> {
44        match self {
45            PrivateSqliteRow::Direct(stmt) => {
46                let column_names = if let Some(column_names) = column_names {
47                    column_names.clone()
48                } else {
49                    let c: Rc<[Option<String>]> = Rc::from(
50                        (0..stmt.column_count())
51                            .map(|idx| stmt.field_name(idx).map(|s| s.to_owned()))
52                            .collect::<Vec<_>>(),
53                    );
54                    *column_names = Some(c.clone());
55                    c
56                };
57                Ok(PrivateSqliteRow::Duplicated {
58                    values: (0..stmt.column_count())
59                        .map(|idx| stmt.copy_value(idx))
60                        .collect::<QueryResult<Vec<_>>>()?,
61                    column_names,
62                })
63            }
64            PrivateSqliteRow::Duplicated {
65                values,
66                column_names,
67            } => Ok(PrivateSqliteRow::Duplicated {
68                values: values
69                    .iter()
70                    .map(|v| v.as_ref().map(OwnedSqliteValue::duplicate).transpose())
71                    .collect::<QueryResult<Vec<_>>>()?,
72                column_names: column_names.clone(),
73            }),
74        }
75    }
76
77    /// Copies the row out of the statement, so that it outlives it.
78    ///
79    /// # Panics
80    ///
81    /// Panics if SQLite cannot allocate the copies, as `IntoOwnedRow::into_owned`
82    /// cannot report an error.
83    pub(super) fn moveable(
84        &self,
85        column_name_cache: &mut Option<Arc<[Option<String>]>>,
86    ) -> OwnedSqliteRow {
87        match self {
88            PrivateSqliteRow::Direct(stmt) => {
89                if column_name_cache.is_none() {
90                    *column_name_cache = Some(
91                        (0..stmt.column_count())
92                            .map(|idx| stmt.field_name(idx).map(|s| s.to_owned()))
93                            .collect::<Vec<_>>()
94                            .into(),
95                    );
96                }
97                let column_names = Arc::clone(
98                    column_name_cache
99                        .as_ref()
100                        .expect("This is initialized above"),
101                );
102                OwnedSqliteRow::new(
103                    (0..stmt.column_count())
104                        .map(|idx| stmt.copy_value(idx))
105                        .collect::<QueryResult<Vec<_>>>()
106                        .unwrap_or_else(|e| { ::core::panicking::panic_fmt(format_args!("{0}", e)); }panic!("{e}")),
107                    column_names,
108                )
109            }
110            PrivateSqliteRow::Duplicated {
111                values,
112                column_names,
113            } => {
114                if column_name_cache.is_none() {
115                    *column_name_cache = Some(
116                        (*column_names)
117                            .iter()
118                            .map(|s| s.to_owned())
119                            .collect::<Vec<_>>()
120                            .into(),
121                    );
122                }
123                let column_names = Arc::clone(
124                    column_name_cache
125                        .as_ref()
126                        .expect("This is initialized above"),
127                );
128                OwnedSqliteRow::new(
129                    values
130                        .iter()
131                        .map(|v| v.as_ref().map(OwnedSqliteValue::duplicate).transpose())
132                        .collect::<QueryResult<Vec<_>>>()
133                        .unwrap_or_else(|e| { ::core::panicking::panic_fmt(format_args!("{0}", e)); }panic!("{e}")),
134                    column_names,
135                )
136            }
137        }
138    }
139}
140
141impl RowSealed for SqliteRow<'_, '_> {}
142
143impl<'stmt> Row<'stmt, Sqlite> for SqliteRow<'stmt, '_> {
144    type Field<'field>
145        = SqliteField<'field, 'field>
146    where
147        'stmt: 'field,
148        Self: 'field;
149    type InnerPartialRow = Self;
150
151    fn field_count(&self) -> usize {
152        self.field_count
153    }
154
155    fn get<'field, I>(&'field self, idx: I) -> Option<Self::Field<'field>>
156    where
157        'stmt: 'field,
158        Self: RowIndex<I>,
159    {
160        let idx = self.idx(idx)?;
161        Some(SqliteField {
162            row: self.inner.borrow(),
163            col_idx: idx,
164        })
165    }
166
167    fn partial_row(&self, range: core::ops::Range<usize>) -> PartialRow<'_, Self::InnerPartialRow> {
168        PartialRow::new(self, range)
169    }
170}
171
172impl RowIndex<usize> for SqliteRow<'_, '_> {
173    fn idx(&self, idx: usize) -> Option<usize> {
174        if idx < self.field_count {
175            Some(idx)
176        } else {
177            None
178        }
179    }
180}
181
182impl<'idx> RowIndex<&'idx str> for SqliteRow<'_, '_> {
183    fn idx(&self, field_name: &'idx str) -> Option<usize> {
184        match &mut *self.inner.borrow_mut() {
185            PrivateSqliteRow::Direct(stmt) => stmt.index_for_column_name(field_name),
186            PrivateSqliteRow::Duplicated { column_names, .. } => column_names
187                .iter()
188                .position(|n| n.as_ref().map(|s| s as &str) == Some(field_name)),
189        }
190    }
191}
192
193#[allow(missing_debug_implementations)]
194pub struct SqliteField<'stmt, 'query> {
195    pub(super) row: Ref<'stmt, PrivateSqliteRow<'stmt, 'query>>,
196    pub(super) col_idx: usize,
197}
198
199impl<'stmt> Field<'stmt, Sqlite> for SqliteField<'stmt, '_> {
200    fn field_name(&self) -> Option<&str> {
201        match &*self.row {
202            PrivateSqliteRow::Direct(stmt) => stmt.field_name(
203                self.col_idx
204                    .try_into()
205                    .expect("Diesel expects to run at least on a 32 bit platform"),
206            ),
207            PrivateSqliteRow::Duplicated { column_names, .. } => column_names
208                .get(self.col_idx)
209                .and_then(|t| t.as_ref().map(|n| n as &str)),
210        }
211    }
212
213    fn is_null(&self) -> bool {
214        self.value().is_none()
215    }
216
217    fn value(&self) -> Option<<Sqlite as Backend>::RawValue<'_>> {
218        SqliteValue::new(Ref::clone(&self.row), self.col_idx)
219    }
220}
221
222// Reads text or blob values, which requires accessing memory allocated by the
223// native library. That is not supported when running under miri with a native
224// libsqlite3 (`-Zmiri-native-lib`), so the whole module is compiled out in
225// that case.
226#[cfg(all(test, not(miri)))]
227mod tests {
228    use super::*;
229
230    #[diesel_test_helper::test]
231    fn fun_with_row_iters() {
232        crate::table! {
233            #[allow(unused_parens)]
234            users(id) {
235                id -> Integer,
236                name -> Text,
237            }
238        }
239
240        use crate::connection::LoadConnection;
241        use crate::deserialize::{FromSql, FromSqlRow};
242        use crate::prelude::*;
243        use crate::row::{Field, Row};
244        use crate::sql_types;
245
246        let conn = &mut crate::test_helpers::connection();
247
248        crate::sql_query("CREATE TABLE users(id INTEGER PRIMARY KEY, name TEXT NOT NULL);")
249            .execute(conn)
250            .unwrap();
251
252        crate::insert_into(users::table)
253            .values(vec![
254                (users::id.eq(1), users::name.eq("Sean")),
255                (users::id.eq(2), users::name.eq("Tess")),
256            ])
257            .execute(conn)
258            .unwrap();
259
260        let query = users::table.select((users::id, users::name));
261
262        let expected = vec![(1, String::from("Sean")), (2, String::from("Tess"))];
263
264        let row_iter = conn.load(query).unwrap();
265        for (row, expected) in row_iter.zip(&expected) {
266            let row = row.unwrap();
267
268            let deserialized = <(i32, String) as FromSqlRow<
269                (sql_types::Integer, sql_types::Text),
270                _,
271            >>::build_from_row(&row)
272            .unwrap();
273
274            assert_eq!(&deserialized, expected);
275        }
276
277        {
278            let collected_rows = conn.load(query).unwrap().collect::<Vec<_>>();
279
280            for (row, expected) in collected_rows.iter().zip(&expected) {
281                let deserialized = row
282                    .as_ref()
283                    .map(|row| {
284                        <(i32, String) as FromSqlRow<
285                            (sql_types::Integer, sql_types::Text),
286                        _,
287                        >>::build_from_row(row).unwrap()
288                    })
289                    .unwrap();
290
291                assert_eq!(&deserialized, expected);
292            }
293        }
294
295        let mut row_iter = conn.load(query).unwrap();
296
297        let first_row = row_iter.next().unwrap().unwrap();
298        let first_fields = (first_row.get(0).unwrap(), first_row.get(1).unwrap());
299        let first_values = (first_fields.0.value(), first_fields.1.value());
300
301        assert!(row_iter.next().unwrap().is_err());
302        std::mem::drop(first_values);
303        assert!(row_iter.next().unwrap().is_err());
304        std::mem::drop(first_fields);
305
306        let second_row = row_iter.next().unwrap().unwrap();
307        let second_fields = (second_row.get(0).unwrap(), second_row.get(1).unwrap());
308        let second_values = (second_fields.0.value(), second_fields.1.value());
309
310        assert!(row_iter.next().unwrap().is_err());
311        std::mem::drop(second_values);
312        assert!(row_iter.next().unwrap().is_err());
313        std::mem::drop(second_fields);
314
315        assert!(row_iter.next().is_none());
316
317        let first_fields = (first_row.get(0).unwrap(), first_row.get(1).unwrap());
318        let second_fields = (second_row.get(0).unwrap(), second_row.get(1).unwrap());
319
320        let first_values = (first_fields.0.value(), first_fields.1.value());
321        let second_values = (second_fields.0.value(), second_fields.1.value());
322
323        assert_eq!(
324            <i32 as FromSql<sql_types::Integer, Sqlite>>::from_nullable_sql(first_values.0)
325                .unwrap(),
326            expected[0].0
327        );
328        assert_eq!(
329            <String as FromSql<sql_types::Text, Sqlite>>::from_nullable_sql(first_values.1)
330                .unwrap(),
331            expected[0].1
332        );
333
334        assert_eq!(
335            <i32 as FromSql<sql_types::Integer, Sqlite>>::from_nullable_sql(second_values.0)
336                .unwrap(),
337            expected[1].0
338        );
339        assert_eq!(
340            <String as FromSql<sql_types::Text, Sqlite>>::from_nullable_sql(second_values.1)
341                .unwrap(),
342            expected[1].1
343        );
344
345        let first_fields = (first_row.get(0).unwrap(), first_row.get(1).unwrap());
346        let first_values = (first_fields.0.value(), first_fields.1.value());
347
348        assert_eq!(
349            <i32 as FromSql<sql_types::Integer, Sqlite>>::from_nullable_sql(first_values.0)
350                .unwrap(),
351            expected[0].0
352        );
353        assert_eq!(
354            <String as FromSql<sql_types::Text, Sqlite>>::from_nullable_sql(first_values.1)
355                .unwrap(),
356            expected[0].1
357        );
358    }
359
360    #[cfg(feature = "returning_clauses_for_sqlite_3_35")]
361    #[crate::declare_sql_function]
362    extern "SQL" {
363        fn sleep(a: diesel::sql_types::Integer) -> diesel::sql_types::Integer;
364    }
365
366    #[diesel_test_helper::test]
367    #[cfg(feature = "returning_clauses_for_sqlite_3_35")]
368    #[allow(clippy::cast_sign_loss)]
369    fn parallel_iter_with_error() {
370        use crate::SqliteConnection;
371        use crate::connection::Connection;
372        use crate::connection::LoadConnection;
373        use crate::connection::SimpleConnection;
374        use crate::expression_methods::ExpressionMethods;
375        use std::sync::{Arc, Barrier};
376        use std::time::Duration;
377
378        let temp_dir = tempfile::tempdir().unwrap();
379        let db_path = format!("{}/test.db", temp_dir.path().display());
380        let mut conn1 = SqliteConnection::establish(&db_path).unwrap();
381        let mut conn2 = SqliteConnection::establish(&db_path).unwrap();
382
383        crate::table! {
384            users {
385                id -> Integer,
386                name -> Text,
387            }
388        }
389
390        conn1
391            .batch_execute("CREATE TABLE users(id INTEGER NOT NULL PRIMARY KEY, name TEXT)")
392            .unwrap();
393
394        let barrier = Arc::new(Barrier::new(2));
395        let barrier2 = barrier.clone();
396
397        // we unblock the main thread from the sleep function
398        sleep_utils::register_nondeterministic_impl(&mut conn2, move |a: i32| {
399            barrier.wait();
400            std::thread::sleep(Duration::from_secs(a as u64));
401            a
402        })
403        .unwrap();
404
405        // spawn a background thread that locks the database file
406        let handle = std::thread::spawn(move || {
407            use crate::query_dsl::RunQueryDsl;
408
409            conn2
410                .immediate_transaction(|conn| diesel::select(sleep(1)).execute(conn))
411                .unwrap();
412        });
413        barrier2.wait();
414
415        // execute some action that also requires a lock
416        let mut iter = conn1
417            .load(
418                diesel::insert_into(users::table)
419                    .values((users::id.eq(1), users::name.eq("John")))
420                    .returning(users::id),
421            )
422            .unwrap();
423
424        // get the first iterator result, that should return the lock error
425        let n = iter.next().unwrap();
426        assert!(n.is_err());
427
428        // check that the iterator is now empty
429        let n = iter.next();
430        assert!(n.is_none());
431
432        // join the background thread
433        handle.join().unwrap();
434    }
435}