Skip to main content

hyperopt_storage/
sqlite.rs

1use hyperopt_core::{Direction, Storage, StorageError, StudyMetadata, Trial};
2use rusqlite::{Connection, OptionalExtension};
3use std::path::Path;
4use std::sync::Mutex;
5
6/// On-disk storage format version. Bumped only on a breaking change to the
7/// SQLite schema; [`SqliteStorage::open`] refuses to read a file written by a
8/// newer, incompatible version rather than silently misreading it.
9const SCHEMA_VERSION: i64 = 1;
10
11/// SQLite-backed storage: trials are persisted to a file so studies survive
12/// process restarts and can be resumed. Mirrors Optuna's RDB storage pattern.
13///
14/// Each trial is stored as a JSON document keyed by `(study_name, number)`,
15/// which keeps the schema stable while faithfully round-tripping the
16/// define-by-run parameter set and every intermediate report.
17///
18/// The connection is guarded by a `Mutex`, so the backend is `Send + Sync` and
19/// usable under parallel execution. Several independent processes can also open
20/// the same file (SQLite handles the file-level locking) for lightweight
21/// multi-process coordination — a practical middle ground short of full
22/// distributed execution.
23pub struct SqliteStorage {
24    conn: Mutex<Connection>,
25}
26
27impl SqliteStorage {
28    /// Open (creating if absent) a study database at `path`.
29    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    /// Open an anonymous in-memory SQLite database (useful for tests).
35    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        // Establish or verify the schema version.
60        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}