Skip to main content

ruda_dataset/dataset/sqlite/
reader.rs

1use std::{
2    marker::PhantomData,
3    path::{Path, PathBuf},
4};
5
6use r2d2::Pool;
7use r2d2_sqlite::{SqliteConnectionManager, rusqlite::OptionalExtension};
8use serde::de::DeserializeOwned;
9use serde_rusqlite::{columns_from_statement, from_row_with_columns};
10
11use super::{Result, SqliteDataset, connection::create_conn_pool};
12use crate::Dataset;
13
14impl<I> SqliteDataset<I> {
15    /// Initializes a `SqliteDataset` from a SQLite database file and a split name.
16    pub fn from_db_file<P: AsRef<Path>>(db_file: P, split: &str) -> Result<Self> {
17        // Create a connection pool
18        let conn_pool = create_conn_pool(&db_file, false)?;
19
20        // Determine how the table is stored
21        let row_serialized = Self::check_if_row_serialized(&conn_pool, split)?;
22
23        // Create a select statement and save it
24        let select_statement = if row_serialized {
25            format!("select item from {split} where row_id = ?")
26        } else {
27            format!("select * from {split} where row_id = ?")
28        };
29
30        // Save the column names and the number of rows
31        let (columns, len) = fetch_columns_and_len(&conn_pool, &select_statement, split)?;
32
33        Ok(SqliteDataset {
34            db_file: db_file.as_ref().to_path_buf(),
35            split: split.to_string(),
36            conn_pool,
37            columns,
38            len,
39            select_statement,
40            row_serialized,
41            phantom: PhantomData,
42        })
43    }
44
45    /// Returns true if table has two columns: row_id (integer) and item (blob).
46    ///
47    /// This is used to determine if the table is row serialized or not.
48    fn check_if_row_serialized(
49        conn_pool: &Pool<SqliteConnectionManager>,
50        split: &str,
51    ) -> Result<bool> {
52        // This struct is used to store the column name and type
53        struct Column {
54            name: String,
55            ty: String,
56        }
57
58        const COLUMN_NAME: usize = 1;
59        const COLUMN_TYPE: usize = 2;
60
61        let sql_statement = format!("PRAGMA table_info({split})");
62
63        let conn = conn_pool.get()?;
64
65        let mut stmt = conn.prepare(sql_statement.as_str())?;
66        let column_iter = stmt.query_map([], |row| {
67            Ok(Column {
68                name: row
69                    .get::<usize, String>(COLUMN_NAME)
70                    .unwrap()
71                    .to_lowercase(),
72                ty: row
73                    .get::<usize, String>(COLUMN_TYPE)
74                    .unwrap()
75                    .to_lowercase(),
76            })
77        })?;
78
79        let mut columns: Vec<Column> = vec![];
80
81        for column in column_iter {
82            columns.push(column?);
83        }
84
85        if columns.len() != 2 {
86            Ok(false)
87        } else {
88            // Check if the column names and types match the expected values
89            Ok(columns[0].name == "row_id"
90                && columns[0].ty == "integer"
91                && columns[1].name == "item"
92                && columns[1].ty == "blob")
93        }
94    }
95
96    /// Get the database file name.
97    pub fn db_file(&self) -> PathBuf {
98        self.db_file.clone()
99    }
100
101    /// Get the split name.
102    pub fn split(&self) -> &str {
103        self.split.as_str()
104    }
105}
106
107impl<I> Dataset<I> for SqliteDataset<I>
108where
109    I: Clone + Send + Sync + DeserializeOwned,
110{
111    /// Get an item from the dataset.
112    fn get(&self, index: usize) -> Option<I> {
113        // Row ids start with 1 (one) and index starts with 0 (zero)
114        let row_id = index + 1;
115
116        // Get a connection from the pool
117        let connection = self.conn_pool.get().unwrap();
118        let mut statement = connection.prepare(self.select_statement.as_str()).unwrap();
119
120        if self.row_serialized {
121            // Fetch with a single column `item` and deserialize it with MessagePack
122            statement
123                .query_row([row_id], |row| {
124                    // Deserialize item (blob) with MessagePack (rmp-serde)
125                    Ok(
126                        rmp_serde::from_slice::<I>(row.get_ref(0).unwrap().as_blob().unwrap())
127                            .unwrap(),
128                    )
129                })
130                .optional() //Converts Error (not found) to None
131                .unwrap()
132        } else {
133            // Fetch a row with multiple columns and deserialize it serde_rusqlite
134            statement
135                .query_row([row_id], |row| {
136                    // Deserialize the row with serde_rusqlite
137                    Ok(from_row_with_columns::<I>(row, &self.columns).unwrap())
138                })
139                .optional() //Converts Error (not found) to None
140                .unwrap()
141        }
142    }
143
144    /// Return the number of rows in the dataset.
145    fn len(&self) -> usize {
146        self.len
147    }
148}
149
150/// Fetch the column names and the number of rows from the database.
151fn fetch_columns_and_len(
152    conn_pool: &Pool<SqliteConnectionManager>,
153    select_statement: &str,
154    split: &str,
155) -> Result<(Vec<String>, usize)> {
156    // Save the column names
157    let connection = conn_pool.get()?;
158    let statement = connection.prepare(select_statement)?;
159    let columns = columns_from_statement(&statement);
160
161    // Count the number of rows and save it as len
162    //
163    // NOTE: Using coalesce(max(row_id), 0) instead of count(*) because count(*) is super slow for large tables.
164    // The coalesce(max(row_id), 0) returns 0 if the table is empty, otherwise it returns the max row_id,
165    // which corresponds to the number of rows in the table.
166    // The main assumption, which always holds true, is that the row_id is always increasing and there are no gaps.
167    // This is true for all the datasets that we are using, otherwise row_id will not correspond to the index.
168    let mut statement =
169        connection.prepare(format!("select coalesce(max(row_id), 0) from {split}").as_str())?;
170
171    let len = statement.query_row([], |row| {
172        let len: usize = row.get(0)?;
173        Ok(len)
174    })?;
175    Ok((columns, len))
176}