ruda_dataset/dataset/sqlite/writer.rs
1use std::{
2 collections::HashSet,
3 fs,
4 marker::PhantomData,
5 path::Path,
6 sync::{Arc, RwLock},
7};
8
9use gix_tempfile::{AutoRemove, ContainingDirectory};
10use r2d2::PooledConnection;
11use r2d2_sqlite::SqliteConnectionManager;
12use serde::{Serialize, de::DeserializeOwned};
13
14use super::{Result, SqliteDatasetError, SqliteDatasetWriter, connection::create_conn_pool};
15
16impl<I> SqliteDatasetWriter<I>
17where
18 I: Clone + Send + Sync + Serialize + DeserializeOwned,
19{
20 /// Creates a new instance of `SqliteDatasetWriter`.
21 ///
22 /// # Arguments
23 ///
24 /// * `db_file` - A reference to the Path that represents the database file path.
25 /// * `overwrite` - A boolean indicating if the existing database file should be overwritten.
26 ///
27 /// # Returns
28 ///
29 /// * A `Result` which is `Ok` if the writer could be created, `Err` otherwise.
30 pub fn new<P: AsRef<Path>>(db_file: P, overwrite: bool) -> Result<Self> {
31 let writer = Self {
32 db_file: db_file.as_ref().to_path_buf(),
33 db_file_tmp: None,
34 splits: Arc::new(RwLock::new(HashSet::new())),
35 overwrite,
36 conn_pool: None,
37 is_completed: Arc::new(RwLock::new(false)),
38 phantom: PhantomData,
39 };
40
41 writer.init()
42 }
43
44 /// Initializes the dataset writer by creating the database file, tables, and connection pool.
45 ///
46 /// # Returns
47 ///
48 /// * A `Result` which is `Ok` if the writer could be initialized, `Err` otherwise.
49 fn init(mut self) -> Result<Self> {
50 // Remove the db file if it already exists
51 if self.db_file.exists() {
52 if self.overwrite {
53 fs::remove_file(&self.db_file)?;
54 } else {
55 return Err(SqliteDatasetError::FileExists(self.db_file));
56 }
57 }
58
59 // Create the database file directory if it does not exist
60 let db_file_dir = self
61 .db_file
62 .parent()
63 .ok_or("Unable to get parent directory")?;
64
65 if !db_file_dir.exists() {
66 fs::create_dir_all(db_file_dir)?;
67 }
68
69 // Create a temp database file name as {base_dir}/{name}.db.tmp
70 let mut db_file_tmp = self.db_file.clone();
71 db_file_tmp.set_extension("db.tmp");
72 if db_file_tmp.exists() {
73 fs::remove_file(&db_file_tmp)?;
74 }
75
76 // Create the temp database file and wrap it with a gix_tempfile::Handle
77 // This will ensure that the temp file is deleted when the writer is dropped
78 // or when process exits with SIGINT or SIGTERM (tempfile crate does not do this)
79 gix_tempfile::signal::setup(Default::default());
80 self.db_file_tmp = Some(gix_tempfile::writable_at(
81 &db_file_tmp,
82 ContainingDirectory::Exists,
83 AutoRemove::Tempfile,
84 )?);
85
86 let conn_pool = create_conn_pool(db_file_tmp, true)?;
87 self.conn_pool = Some(conn_pool);
88
89 Ok(self)
90 }
91
92 /// Serializes and writes an item to the database. The item is written to the table for the
93 /// specified split. If the table does not exist, it is created. If the table exists, the item
94 /// is appended to the table. The serialization is done using the [MessagePack](https://msgpack.org/)
95 ///
96 /// # Arguments
97 ///
98 /// * `split` - A string slice that defines the data split for writing (e.g., "train", "test").
99 /// * `item` - A reference to the item to be written to the database.
100 ///
101 /// # Returns
102 ///
103 /// * A `Result` containing the index of the inserted row if successful, an error otherwise.
104 pub fn write(&self, split: &str, item: &I) -> Result<usize> {
105 // Acquire the read lock (wont't block other reads)
106 let is_completed = self.is_completed.read().unwrap();
107
108 // If the writer is completed, return an error
109 if *is_completed {
110 return Err(SqliteDatasetError::Other(
111 "Cannot save to a completed dataset writer",
112 ));
113 }
114
115 // create the table for the split if it does not exist
116 if !self.splits.read().unwrap().contains(split) {
117 self.create_table(split)?;
118 }
119
120 // Get a connection from the pool
121 let conn_pool = self.conn_pool.as_ref().unwrap();
122 let conn = conn_pool.get()?;
123
124 // Serialize the item using MessagePack
125 let serialized_item = rmp_serde::to_vec(item)?;
126
127 // Turn off the synchronous and journal mode for speed up
128 // We are sacrificing durability for speed but it's okay because
129 // we always recreate the dataset if it is not completed.
130 pragma_update_with_error_handling(&conn, "synchronous", "OFF")?;
131 pragma_update_with_error_handling(&conn, "journal_mode", "OFF")?;
132
133 // Insert the serialized item into the database
134 let insert_statement = format!("insert into {split} (item) values (?)");
135 conn.execute(insert_statement.as_str(), [serialized_item])?;
136
137 // Get the primary key of the last inserted row and convert to index (row_id-1)
138 let index = (conn.last_insert_rowid() - 1) as usize;
139
140 Ok(index)
141 }
142
143 /// Marks the dataset as completed and persists the temporary database file.
144 pub fn set_completed(&mut self) -> Result<()> {
145 let mut is_completed = self.is_completed.write().unwrap();
146
147 // Force close the connection pool
148 // This is required on Windows platform where the connection pool prevents
149 // from persisting the db by renaming the temp file.
150 if let Some(pool) = self.conn_pool.take() {
151 std::mem::drop(pool);
152 }
153
154 // Rename the database file from tmp to db
155 let _file_result = self
156 .db_file_tmp
157 .take() // take ownership of the temporary file and set to None
158 .unwrap() // unwrap the temporary file
159 .persist(&self.db_file)?
160 .ok_or("Unable to persist the database file")?;
161
162 *is_completed = true;
163 Ok(())
164 }
165
166 /// Creates table for the data split.
167 ///
168 /// Note: call is idempotent and thread-safe.
169 ///
170 /// # Arguments
171 ///
172 /// * `split` - A string slice that defines the data split for the table (e.g., "train", "test").
173 ///
174 /// # Returns
175 ///
176 /// * A `Result` which is `Ok` if the table could be created, `Err` otherwise.
177 ///
178 /// TODO (@antimora): add support creating a table with columns corresponding to the item fields
179 fn create_table(&self, split: &str) -> Result<()> {
180 // Check if the split already exists
181 if self.splits.read().unwrap().contains(split) {
182 return Ok(());
183 }
184
185 let conn_pool = self.conn_pool.as_ref().unwrap();
186 let connection = conn_pool.get()?;
187 let create_table_statement = format!(
188 "create table if not exists {split} (row_id integer primary key autoincrement not \
189 null, item blob not null)"
190 );
191
192 connection.execute(create_table_statement.as_str(), [])?;
193
194 // Add the split to the splits
195 self.splits.write().unwrap().insert(split.to_string());
196
197 Ok(())
198 }
199}
200
201/// Runs a pragma update and ignores the `ExecuteReturnedResults` error.
202///
203/// Sometimes ExecuteReturnedResults is returned when running a pragma update. This is not an error
204/// and can be ignored. This function runs the pragma update and ignores the error if it is
205/// `ExecuteReturnedResults`.
206fn pragma_update_with_error_handling(
207 conn: &PooledConnection<SqliteConnectionManager>,
208 setting: &str,
209 value: &str,
210) -> Result<()> {
211 let result = conn.pragma_update(None, setting, value);
212 if let Err(error) = result
213 && error != rusqlite::Error::ExecuteReturnedResults
214 {
215 return Err(SqliteDatasetError::Sql(error));
216 }
217
218 Ok(())
219}