ruda_dataset/dataset/sqlite/
reader.rs1use 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 pub fn from_db_file<P: AsRef<Path>>(db_file: P, split: &str) -> Result<Self> {
17 let conn_pool = create_conn_pool(&db_file, false)?;
19
20 let row_serialized = Self::check_if_row_serialized(&conn_pool, split)?;
22
23 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 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 fn check_if_row_serialized(
49 conn_pool: &Pool<SqliteConnectionManager>,
50 split: &str,
51 ) -> Result<bool> {
52 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 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 pub fn db_file(&self) -> PathBuf {
98 self.db_file.clone()
99 }
100
101 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 fn get(&self, index: usize) -> Option<I> {
113 let row_id = index + 1;
115
116 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 statement
123 .query_row([row_id], |row| {
124 Ok(
126 rmp_serde::from_slice::<I>(row.get_ref(0).unwrap().as_blob().unwrap())
127 .unwrap(),
128 )
129 })
130 .optional() .unwrap()
132 } else {
133 statement
135 .query_row([row_id], |row| {
136 Ok(from_row_with_columns::<I>(row, &self.columns).unwrap())
138 })
139 .optional() .unwrap()
141 }
142 }
143
144 fn len(&self) -> usize {
146 self.len
147 }
148}
149
150fn fetch_columns_and_len(
152 conn_pool: &Pool<SqliteConnectionManager>,
153 select_statement: &str,
154 split: &str,
155) -> Result<(Vec<String>, usize)> {
156 let connection = conn_pool.get()?;
158 let statement = connection.prepare(select_statement)?;
159 let columns = columns_from_statement(&statement);
160
161 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}