Skip to main content

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}