Skip to main content

ruda_dataset/dataset/sqlite/
storage.rs

1use std::path::{Path, PathBuf};
2
3use sanitize_filename::sanitize;
4use serde::{Serialize, de::DeserializeOwned};
5
6use super::{Result, SqliteDataset, SqliteDatasetStorage, SqliteDatasetWriter};
7
8impl SqliteDatasetStorage {
9    /// Creates a new instance of `SqliteDatasetStorage` using a dataset name.
10    ///
11    /// # Arguments
12    ///
13    /// * `name` - A string slice that holds the name of the dataset.
14    pub fn from_name(name: &str) -> Self {
15        SqliteDatasetStorage {
16            name: Some(name.to_string()),
17            db_file: None,
18            base_dir: None,
19        }
20    }
21
22    /// Creates a new instance of `SqliteDatasetStorage` using a database file path.
23    ///
24    /// # Arguments
25    ///
26    /// * `db_file` - A reference to the Path that represents the database file path.
27    pub fn from_file<P: AsRef<Path>>(db_file: P) -> Self {
28        SqliteDatasetStorage {
29            name: None,
30            db_file: Some(db_file.as_ref().to_path_buf()),
31            base_dir: None,
32        }
33    }
34
35    /// Sets the base directory for storing the dataset.
36    ///
37    /// # Arguments
38    ///
39    /// * `base_dir` - A string slice that represents the base directory.
40    pub fn with_base_dir<P: AsRef<Path>>(mut self, base_dir: P) -> Self {
41        self.base_dir = Some(base_dir.as_ref().to_path_buf());
42        self
43    }
44
45    /// Checks if the database file exists in the given path.
46    ///
47    /// # Returns
48    ///
49    /// * A boolean value indicating whether the file exists or not.
50    pub fn exists(&self) -> bool {
51        self.db_file().exists()
52    }
53
54    /// Fetches the database file path.
55    ///
56    /// # Returns
57    ///
58    /// * A `PathBuf` instance representing the file path.
59    pub fn db_file(&self) -> PathBuf {
60        match &self.db_file {
61            Some(db_file) => db_file.clone(),
62            None => {
63                let name = sanitize(self.name.as_ref().expect("Name is not set"));
64                Self::base_dir(self.base_dir.to_owned()).join(format!("{name}.db"))
65            }
66        }
67    }
68
69    /// Determines the base directory for storing the dataset.
70    ///
71    /// # Arguments
72    ///
73    /// * `base_dir` - An `Option` that may contain a `PathBuf` instance representing the base directory.
74    ///
75    /// # Returns
76    ///
77    /// * A `PathBuf` instance representing the base directory.
78    pub fn base_dir(base_dir: Option<PathBuf>) -> PathBuf {
79        match base_dir {
80            Some(base_dir) => base_dir,
81            None => dirs::cache_dir()
82                .expect("Could not get cache directory")
83                .join("ruda-dataset"),
84        }
85    }
86
87    /// Provides a writer instance for the SQLite dataset.
88    ///
89    /// # Arguments
90    ///
91    /// * `overwrite` - A boolean indicating if the existing database file should be overwritten.
92    ///
93    /// # Returns
94    ///
95    /// * A `Result` which is `Ok` if the writer could be created, `Err` otherwise.
96    pub fn writer<I>(&self, overwrite: bool) -> Result<SqliteDatasetWriter<I>>
97    where
98        I: Clone + Send + Sync + Serialize + DeserializeOwned,
99    {
100        SqliteDatasetWriter::new(self.db_file(), overwrite)
101    }
102
103    /// Provides a reader instance for the SQLite dataset.
104    ///
105    /// # Arguments
106    ///
107    /// * `split` - A string slice that defines the data split for reading (e.g., "train", "test").
108    ///
109    /// # Returns
110    ///
111    /// * A `Result` which is `Ok` if the reader could be created, `Err` otherwise.
112    pub fn reader<I>(&self, split: &str) -> Result<SqliteDataset<I>>
113    where
114        I: Clone + Send + Sync + Serialize + DeserializeOwned,
115    {
116        if !self.exists() {
117            panic!("The database file does not exist");
118        }
119
120        SqliteDataset::from_db_file(self.db_file(), split)
121    }
122}