hyperopt_storage/
sqlite.rs1use hyperopt_core::{Direction, Storage, StorageError, StudyMetadata, Trial};
2use rusqlite::{Connection, OptionalExtension};
3use std::path::Path;
4use std::sync::Mutex;
5
6const SCHEMA_VERSION: i64 = 1;
10
11pub struct SqliteStorage {
24 conn: Mutex<Connection>,
25}
26
27impl SqliteStorage {
28 pub fn open(path: impl AsRef<Path>) -> Result<Self, StorageError> {
30 let conn = Connection::open(path).map_err(backend)?;
31 Self::from_connection(conn)
32 }
33
34 pub fn open_in_memory() -> Result<Self, StorageError> {
36 let conn = Connection::open_in_memory().map_err(backend)?;
37 Self::from_connection(conn)
38 }
39
40 fn from_connection(conn: Connection) -> Result<Self, StorageError> {
41 conn.execute_batch(
42 "CREATE TABLE IF NOT EXISTS hyperopt_meta (
43 id INTEGER PRIMARY KEY CHECK (id = 1),
44 schema_version INTEGER NOT NULL
45 );
46 CREATE TABLE IF NOT EXISTS studies (
47 name TEXT PRIMARY KEY,
48 direction TEXT NOT NULL
49 );
50 CREATE TABLE IF NOT EXISTS trials (
51 study_name TEXT NOT NULL,
52 number INTEGER NOT NULL,
53 data TEXT NOT NULL,
54 PRIMARY KEY (study_name, number)
55 );",
56 )
57 .map_err(backend)?;
58
59 let existing: Option<i64> = conn
61 .query_row(
62 "SELECT schema_version FROM hyperopt_meta WHERE id = 1",
63 [],
64 |row| row.get(0),
65 )
66 .optional()
67 .map_err(backend)?;
68 match existing {
69 None => {
70 conn.execute(
71 "INSERT INTO hyperopt_meta (id, schema_version) VALUES (1, ?1)",
72 [SCHEMA_VERSION],
73 )
74 .map_err(backend)?;
75 }
76 Some(v) if v != SCHEMA_VERSION => {
77 return Err(StorageError::SchemaMismatch {
78 found: v,
79 expected: SCHEMA_VERSION,
80 });
81 }
82 Some(_) => {}
83 }
84
85 Ok(SqliteStorage {
86 conn: Mutex::new(conn),
87 })
88 }
89
90 fn lock(&self) -> std::sync::MutexGuard<'_, Connection> {
91 self.conn.lock().unwrap_or_else(|p| p.into_inner())
92 }
93}
94
95impl Storage for SqliteStorage {
96 fn save_trial(&self, study_name: &str, trial: &Trial) -> Result<(), StorageError> {
97 let data = serde_json::to_string(trial)
98 .map_err(|e| StorageError::Serialization(e.to_string()))?;
99 let conn = self.lock();
100 conn.execute(
101 "INSERT INTO trials (study_name, number, data) VALUES (?1, ?2, ?3)
102 ON CONFLICT(study_name, number) DO UPDATE SET data = excluded.data",
103 rusqlite::params![study_name, trial.number as i64, data],
104 )
105 .map_err(backend)?;
106 Ok(())
107 }
108
109 fn load_trials(&self, study_name: &str) -> Result<Vec<Trial>, StorageError> {
110 let conn = self.lock();
111 let mut stmt = conn
112 .prepare("SELECT data FROM trials WHERE study_name = ?1 ORDER BY number ASC")
113 .map_err(backend)?;
114 let rows = stmt
115 .query_map([study_name], |row| row.get::<_, String>(0))
116 .map_err(backend)?;
117 let mut trials = Vec::new();
118 for row in rows {
119 let data = row.map_err(backend)?;
120 let trial: Trial = serde_json::from_str(&data)
121 .map_err(|e| StorageError::Serialization(e.to_string()))?;
122 trials.push(trial);
123 }
124 Ok(trials)
125 }
126
127 fn save_study_metadata(&self, meta: &StudyMetadata) -> Result<(), StorageError> {
128 let conn = self.lock();
129 conn.execute(
130 "INSERT INTO studies (name, direction) VALUES (?1, ?2)
131 ON CONFLICT(name) DO UPDATE SET direction = excluded.direction",
132 rusqlite::params![meta.study_name, direction_to_str(meta.direction)],
133 )
134 .map_err(backend)?;
135 Ok(())
136 }
137
138 fn load_study_metadata(
139 &self,
140 study_name: &str,
141 ) -> Result<Option<StudyMetadata>, StorageError> {
142 let conn = self.lock();
143 let row: Option<String> = conn
144 .query_row(
145 "SELECT direction FROM studies WHERE name = ?1",
146 [study_name],
147 |row| row.get(0),
148 )
149 .optional()
150 .map_err(backend)?;
151 match row {
152 Some(dir) => Ok(Some(StudyMetadata {
153 study_name: study_name.to_string(),
154 direction: direction_from_str(&dir)?,
155 })),
156 None => Ok(None),
157 }
158 }
159}
160
161fn backend(e: rusqlite::Error) -> StorageError {
162 StorageError::Backend(e.to_string())
163}
164
165fn direction_to_str(d: Direction) -> &'static str {
166 match d {
167 Direction::Minimize => "minimize",
168 Direction::Maximize => "maximize",
169 }
170}
171
172fn direction_from_str(s: &str) -> Result<Direction, StorageError> {
173 match s {
174 "minimize" => Ok(Direction::Minimize),
175 "maximize" => Ok(Direction::Maximize),
176 other => Err(StorageError::Backend(format!("unknown direction: {other}"))),
177 }
178}