1#![warn(clippy::all)]
53#![forbid(unsafe_code)]
54
55use std::collections::HashSet;
56use std::error::Error;
57use std::marker::{PhantomData, Send, Sync};
58
59use rusqlite::{params, Connection, Error as RusqliteError, Transaction};
60use uuid::Uuid;
61
62use schemerz::{Adapter, Migration};
63
64pub trait RusqliteMigration: Migration<Uuid> {
66 type Error: From<RusqliteError>;
67
68 fn up(&self, _transaction: &Transaction<'_>) -> Result<(), Self::Error> {
70 Ok(())
71 }
72
73 fn down(&self, _transaction: &Transaction<'_>) -> Result<(), Self::Error> {
75 Ok(())
76 }
77}
78
79pub type RusqliteAdapterError = RusqliteError;
80
81struct WrappedUuid(Uuid);
82
83impl rusqlite::types::FromSql for WrappedUuid {
84 fn column_result(value: rusqlite::types::ValueRef<'_>) -> rusqlite::types::FromSqlResult<Self> {
85 Ok(WrappedUuid(Uuid::from_slice(value.as_blob()?).map_err(
86 |e| rusqlite::types::FromSqlError::Other(Box::new(e)),
87 )?))
88 }
89}
90
91pub struct RusqliteAdapter<'a, E> {
93 conn: &'a mut Connection,
94 migration_metadata_table: String,
95 _err: PhantomData<E>,
96}
97
98impl<'a, E> RusqliteAdapter<'a, E> {
99 pub fn new(conn: &'a mut Connection, table_name: Option<String>) -> RusqliteAdapter<'a, E> {
115 RusqliteAdapter {
116 conn,
117 migration_metadata_table: table_name.unwrap_or_else(|| "_schemerz".into()),
118 _err: PhantomData,
119 }
120 }
121
122 pub fn init(&self) -> Result<(), RusqliteError> {
125 self.conn.execute(
126 &format!(
127 r#"
128 CREATE TABLE IF NOT EXISTS {} (
129 id blob PRIMARY KEY
130 )
131 "#,
132 self.migration_metadata_table
133 ),
134 params![],
135 )?;
136 Ok(())
137 }
138}
139
140impl<'a, E> Adapter<Uuid> for RusqliteAdapter<'a, E>
141where
142 E: From<RusqliteError> + Sync + Send + Error + 'static,
143{
144 type MigrationType = Box<dyn RusqliteMigration<Error = E>>;
145
146 type Error = E;
147
148 fn applied_migrations(&mut self) -> Result<HashSet<Uuid>, Self::Error> {
149 let mut stmt = self.conn.prepare(&format!(
150 "SELECT id FROM {};",
151 self.migration_metadata_table
152 ))?;
153 let rows = stmt.query_map(params![], |row| row.get::<_, WrappedUuid>(0))?;
156 let mut ids = HashSet::new();
157 for row in rows {
158 ids.insert(row?.0);
159 }
160 Ok(ids)
161 }
162
163 fn apply_migration(&mut self, migration: &Self::MigrationType) -> Result<(), Self::Error> {
164 let trans = self.conn.transaction()?;
165 migration.up(&trans)?;
166 let uuid = migration.id();
167 let uuid_bytes = &uuid.as_bytes()[..];
168 trans.execute(
169 &format!(
170 "INSERT INTO {} (id) VALUES (?1);",
171 self.migration_metadata_table
172 ),
173 [&uuid_bytes],
174 )?;
175 trans.commit().map_err(|e| e.into())
176 }
177
178 fn revert_migration(&mut self, migration: &Self::MigrationType) -> Result<(), Self::Error> {
179 let trans = self.conn.transaction()?;
180 migration.down(&trans)?;
181 let uuid = migration.id();
182 let uuid_bytes = &uuid.as_bytes()[..];
183 trans.execute(
184 &format!(
185 "DELETE FROM {} WHERE id = ?1;",
186 self.migration_metadata_table
187 ),
188 [&uuid_bytes],
189 )?;
190 trans.commit().map_err(|e| e.into())
191 }
192}
193
194#[cfg(test)]
195mod tests {
196 use super::*;
197 use rusqlite::Error as RusqliteError;
198 use schemerz::test_schemerz_adapter;
199 use schemerz::testing::*;
200
201 impl RusqliteMigration for TestMigration<Uuid> {
202 type Error = RusqliteError;
203 }
204
205 impl<'a> TestAdapter<Uuid> for RusqliteAdapter<'a, RusqliteError> {
206 fn mock(id: Uuid, dependencies: HashSet<Uuid>) -> Self::MigrationType {
207 Box::new(TestMigration::new(id, dependencies))
208 }
209 }
210
211 fn build_test_connection() -> Connection {
212 Connection::open_in_memory().unwrap()
213 }
214
215 fn build_test_adapter(conn: &mut Connection) -> RusqliteAdapter<'_, RusqliteError> {
216 let adapter = RusqliteAdapter::new(conn, None);
217 adapter.init().unwrap();
218 adapter
219 }
220
221 fn uuid_iter() -> impl Iterator<Item = Uuid> {
222 (0..).map(|v| Uuid::from_fields(v as u32, v, v, &[0; 8]))
223 }
224
225 test_schemerz_adapter!(
226 let mut conn = build_test_connection(),
227 build_test_adapter(&mut conn),
228 uuid_iter(),
229 );
230}