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#[cfg(test)]
223mod tests {
224    use super::*;
225
226    #[diesel_test_helper::test]
227    fn fun_with_row_iters() {
228        crate::table! {
229            #[allow(unused_parens)]
230            users(id) {
231                id -> Integer,
232                name -> Text,
233            }
234        }
235
236        use crate::connection::LoadConnection;
237        use crate::deserialize::{FromSql, FromSqlRow};
238        use crate::prelude::*;
239        use crate::row::{Field, Row};
240        use crate::sql_types;
241
242        let conn = &mut crate::test_helpers::connection();
243
244        crate::sql_query("CREATE TABLE users(id INTEGER PRIMARY KEY, name TEXT NOT NULL);")
245            .execute(conn)
246            .unwrap();
247
248        crate::insert_into(users::table)
249            .values(vec![
250                (users::id.eq(1), users::name.eq("Sean")),
251                (users::id.eq(2), users::name.eq("Tess")),
252            ])
253            .execute(conn)
254            .unwrap();
255
256        let query = users::table.select((users::id, users::name));
257
258        let expected = vec![(1, String::from("Sean")), (2, String::from("Tess"))];
259
260        let row_iter = conn.load(query).unwrap();
261        for (row, expected) in row_iter.zip(&expected) {
262            let row = row.unwrap();
263
264            let deserialized = <(i32, String) as FromSqlRow<
265                (sql_types::Integer, sql_types::Text),
266                _,
267            >>::build_from_row(&row)
268            .unwrap();
269
270            assert_eq!(&deserialized, expected);
271        }
272
273        {
274            let collected_rows = conn.load(query).unwrap().collect::<Vec<_>>();
275
276            for (row, expected) in collected_rows.iter().zip(&expected) {
277                let deserialized = row
278                    .as_ref()
279                    .map(|row| {
280                        <(i32, String) as FromSqlRow<
281                            (sql_types::Integer, sql_types::Text),
282                        _,
283                        >>::build_from_row(row).unwrap()
284                    })
285                    .unwrap();
286
287                assert_eq!(&deserialized, expected);
288            }
289        }
290
291        let mut row_iter = conn.load(query).unwrap();
292
293        let first_row = row_iter.next().unwrap().unwrap();
294        let first_fields = (first_row.get(0).unwrap(), first_row.get(1).unwrap());
295        let first_values = (first_fields.0.value(), first_fields.1.value());
296
297        assert!(row_iter.next().unwrap().is_err());
298        std::mem::drop(first_values);
299        assert!(row_iter.next().unwrap().is_err());
300        std::mem::drop(first_fields);
301
302        let second_row = row_iter.next().unwrap().unwrap();
303        let second_fields = (second_row.get(0).unwrap(), second_row.get(1).unwrap());
304        let second_values = (second_fields.0.value(), second_fields.1.value());
305
306        assert!(row_iter.next().unwrap().is_err());
307        std::mem::drop(second_values);
308        assert!(row_iter.next().unwrap().is_err());
309        std::mem::drop(second_fields);
310
311        assert!(row_iter.next().is_none());
312
313        let first_fields = (first_row.get(0).unwrap(), first_row.get(1).unwrap());
314        let second_fields = (second_row.get(0).unwrap(), second_row.get(1).unwrap());
315
316        let first_values = (first_fields.0.value(), first_fields.1.value());
317        let second_values = (second_fields.0.value(), second_fields.1.value());
318
319        assert_eq!(
320            <i32 as FromSql<sql_types::Integer, Sqlite>>::from_nullable_sql(first_values.0)
321                .unwrap(),
322            expected[0].0
323        );
324        assert_eq!(
325            <String as FromSql<sql_types::Text, Sqlite>>::from_nullable_sql(first_values.1)
326                .unwrap(),
327            expected[0].1
328        );
329
330        assert_eq!(
331            <i32 as FromSql<sql_types::Integer, Sqlite>>::from_nullable_sql(second_values.0)
332                .unwrap(),
333            expected[1].0
334        );
335        assert_eq!(
336            <String as FromSql<sql_types::Text, Sqlite>>::from_nullable_sql(second_values.1)
337                .unwrap(),
338            expected[1].1
339        );
340
341        let first_fields = (first_row.get(0).unwrap(), first_row.get(1).unwrap());
342        let first_values = (first_fields.0.value(), first_fields.1.value());
343
344        assert_eq!(
345            <i32 as FromSql<sql_types::Integer, Sqlite>>::from_nullable_sql(first_values.0)
346                .unwrap(),
347            expected[0].0
348        );
349        assert_eq!(
350            <String as FromSql<sql_types::Text, Sqlite>>::from_nullable_sql(first_values.1)
351                .unwrap(),
352            expected[0].1
353        );
354    }
355
356    #[cfg(feature = "returning_clauses_for_sqlite_3_35")]
357    #[crate::declare_sql_function]
358    extern "SQL" {
359        fn sleep(a: diesel::sql_types::Integer) -> diesel::sql_types::Integer;
360    }
361
362    #[diesel_test_helper::test]
363    #[cfg(feature = "returning_clauses_for_sqlite_3_35")]
364    #[allow(clippy::cast_sign_loss)]
365    fn parallel_iter_with_error() {
366        use crate::SqliteConnection;
367        use crate::connection::Connection;
368        use crate::connection::LoadConnection;
369        use crate::connection::SimpleConnection;
370        use crate::expression_methods::ExpressionMethods;
371        use std::sync::{Arc, Barrier};
372        use std::time::Duration;
373
374        let temp_dir = tempfile::tempdir().unwrap();
375        let db_path = format!("{}/test.db", temp_dir.path().display());
376        let mut conn1 = SqliteConnection::establish(&db_path).unwrap();
377        let mut conn2 = SqliteConnection::establish(&db_path).unwrap();
378
379        crate::table! {
380            users {
381                id -> Integer,
382                name -> Text,
383            }
384        }
385
386        conn1
387            .batch_execute("CREATE TABLE users(id INTEGER NOT NULL PRIMARY KEY, name TEXT)")
388            .unwrap();
389
390        let barrier = Arc::new(Barrier::new(2));
391        let barrier2 = barrier.clone();
392
393        // we unblock the main thread from the sleep function
394        sleep_utils::register_nondeterministic_impl(&mut conn2, move |a: i32| {
395            barrier.wait();
396            std::thread::sleep(Duration::from_secs(a as u64));
397            a
398        })
399        .unwrap();
400
401        // spawn a background thread that locks the database file
402        let handle = std::thread::spawn(move || {
403            use crate::query_dsl::RunQueryDsl;
404
405            conn2
406                .immediate_transaction(|conn| diesel::select(sleep(1)).execute(conn))
407                .unwrap();
408        });
409        barrier2.wait();
410
411        // execute some action that also requires a lock
412        let mut iter = conn1
413            .load(
414                diesel::insert_into(users::table)
415                    .values((users::id.eq(1), users::name.eq("John")))
416                    .returning(users::id),
417            )
418            .unwrap();
419
420        // get the first iterator result, that should return the lock error
421        let n = iter.next().unwrap();
422        assert!(n.is_err());
423
424        // check that the iterator is now empty
425        let n = iter.next();
426        assert!(n.is_none());
427
428        // join the background thread
429        handle.join().unwrap();
430    }
431}