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}