Skip to main content

diesel/sqlite/connection/
row.rs

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