Skip to main content

ruda_dataset/dataset/
sqlite.rs

1mod connection;
2mod reader;
3mod storage;
4mod writer;
5
6#[cfg(test)]
7mod tests;
8
9use std::{
10    collections::HashSet,
11    io,
12    marker::PhantomData,
13    path::PathBuf,
14    sync::{Arc, RwLock},
15};
16
17use gix_tempfile::{
18    Handle,
19    handle::{Writable, persist},
20};
21use r2d2::Pool;
22use r2d2_sqlite::SqliteConnectionManager;
23
24/// Result type for the sqlite dataset.
25pub type Result<T> = core::result::Result<T, SqliteDatasetError>;
26
27/// Sqlite dataset error.
28#[derive(thiserror::Error, Debug)]
29pub enum SqliteDatasetError {
30    /// IO related error.
31    #[error("IO error: {0}")]
32    Io(#[from] io::Error),
33
34    /// Sql related error.
35    #[error("Sql error: {0}")]
36    Sql(#[from] serde_rusqlite::rusqlite::Error),
37
38    /// Serde related error.
39    #[error("Serde error: {0}")]
40    Serde(#[from] rmp_serde::encode::Error),
41
42    /// The database file already exists error.
43    #[error("Overwrite flag is set to false and the database file already exists: {0}")]
44    FileExists(PathBuf),
45
46    /// Error when creating the connection pool.
47    #[error("Failed to create connection pool: {0}")]
48    ConnectionPool(#[from] r2d2::Error),
49
50    /// Error when persisting the temporary database file.
51    #[error("Could not persist the temporary database file: {0}")]
52    PersistDbFile(#[from] persist::Error<Writable>),
53
54    /// Any other error.
55    #[error("{0}")]
56    Other(&'static str),
57}
58
59impl From<&'static str> for SqliteDatasetError {
60    fn from(s: &'static str) -> Self {
61        SqliteDatasetError::Other(s)
62    }
63}
64
65/// This struct represents a dataset where all items are stored in an SQLite database.
66/// Each instance of this struct corresponds to a specific table within the SQLite database,
67/// and allows for interaction with the data stored in the table in a structured and typed manner.
68///
69/// The SQLite database must contain a table with the same name as the `split` field. This table should
70/// have a primary key column named `row_id`, which is used to index the rows in the table. The `row_id`
71/// should start at 1, while the corresponding dataset `index` should start at 0, i.e., `row_id` = `index` + 1.
72///
73/// Table columns can be represented in two ways:
74///
75/// 1. The table can have a column for each field in the `I` struct. In this case, the column names in the table
76///    should match the field names of the `I` struct. The field names can be a subset of column names and
77///    can be in any order.
78///
79/// For the supported field types, refer to:
80/// - [Serialization field types](https://docs.rs/serde_rusqlite/latest/serde_rusqlite)
81/// - [SQLite data types](https://www.sqlite.org/datatype3.html)
82///
83/// 2. The fields in the `I` struct can be serialized into a single column `item` in the table. In this case, the table
84///    should have a single column named `item` of type `BLOB`. This is useful when the `I` struct contains complex fields
85///    that cannot be mapped to a SQLite type, such as nested structs, vectors, etc. The serialization is done using
86///    [MessagePack](https://msgpack.org/).
87///
88/// Note: The code automatically figures out which of the above two cases is applicable, and uses the appropriate
89/// method to read the data from the table.
90#[derive(Debug)]
91pub struct SqliteDataset<I> {
92    db_file: PathBuf,
93    split: String,
94    conn_pool: Pool<SqliteConnectionManager>,
95    columns: Vec<String>,
96    len: usize,
97    select_statement: String,
98    row_serialized: bool,
99    phantom: PhantomData<I>,
100}
101
102/// The `SqliteDatasetStorage` struct represents a SQLite database for storing datasets.
103/// It consists of an optional name, a database file path, and a base directory for storage.
104#[derive(Clone, Debug)]
105pub struct SqliteDatasetStorage {
106    name: Option<String>,
107    db_file: Option<PathBuf>,
108    base_dir: Option<PathBuf>,
109}
110
111/// This `SqliteDatasetWriter` struct is a SQLite database writer dedicated to storing datasets.
112/// It retains the current writer's state and its database connection.
113///
114/// Being thread-safe, this writer can be concurrently used across multiple threads.
115///
116/// Typical applications include:
117///
118/// - Generation of a new dataset
119/// - Storage of preprocessed data or metadata
120/// - Enlargement of a dataset's item count post preprocessing
121#[derive(Debug)]
122pub struct SqliteDatasetWriter<I> {
123    db_file: PathBuf,
124    db_file_tmp: Option<Handle<Writable>>,
125    splits: Arc<RwLock<HashSet<String>>>,
126    overwrite: bool,
127    conn_pool: Option<Pool<SqliteConnectionManager>>,
128    is_completed: Arc<RwLock<bool>>,
129    phantom: PhantomData<I>,
130}