diff options
Diffstat (limited to 'src/main/rust/store')
| -rw-r--r-- | src/main/rust/store/mod.rs | 34 | ||||
| -rw-r--r-- | src/main/rust/store/sqlite.rs | 163 |
2 files changed, 188 insertions, 9 deletions
diff --git a/src/main/rust/store/mod.rs b/src/main/rust/store/mod.rs index 42f1809..7c7d9ce 100644 --- a/src/main/rust/store/mod.rs +++ b/src/main/rust/store/mod.rs @@ -1,11 +1,23 @@ -pub mod sqlite; +mod sqlite; -use miette::Result; +use std::num::TryFromIntError; + +use ::sqlite::Error as SQLiteLibraryError; +use miette::{Diagnostic, Result}; +use thiserror::Error; pub enum Context { SQLite(sqlite::SQLite), } +#[derive(Debug, Diagnostic, Error)] +pub enum ContextError { + #[error("SQLite error")] + SQLiteError(#[from] SQLiteLibraryError), + #[error("Int conversion error")] + ConversionError(#[from] TryFromIntError), +} + impl Default for Context { fn default() -> Self { Context::SQLite(sqlite::SQLite::default()) @@ -13,7 +25,23 @@ impl Default for Context { } pub trait Store { - fn init(&mut self) -> Result<()>; + fn init(&mut self) -> Result<(), ContextError>; + fn get_counter(&self) -> Result<u8, ContextError>; + fn set_counter(&mut self, new_count: u8) -> Result<u8, ContextError>; +} + +impl Context { + pub fn unwrap_store(&mut self) -> &mut impl Store { + match self { + Context::SQLite(store) => store, + } + } + + pub fn get_store(&self) -> &impl Store { + match self { + Context::SQLite(store) => store, + } + } } #[cfg(test)] diff --git a/src/main/rust/store/sqlite.rs b/src/main/rust/store/sqlite.rs index 734c635..6033494 100644 --- a/src/main/rust/store/sqlite.rs +++ b/src/main/rust/store/sqlite.rs @@ -1,5 +1,5 @@ -use crate::store::Store; -use miette::{IntoDiagnostic, Result}; +use crate::store::{ContextError, Store}; +use miette::Result; use sqlite::Connection; pub struct SQLite { @@ -7,6 +7,49 @@ pub struct SQLite { connection: Option<Connection>, } +impl SQLite { + /// Wrap connection in a Result to make it easier to unwrap later + fn get_connection(&self) -> Result<&Connection, sqlite::Error> { + match &self.connection { + Some(conn) => Ok(conn), + None => Err(sqlite::Error { + code: None, + message: Some("Uninitialized".into()), + }), + } + } + + fn create_table(&self) -> Result<(), sqlite::Error> { + let conn = self.get_connection()?; + let mut statement = conn.prepare("CREATE TABLE RATAT ( COUNTER SMALLINT )")?; + statement.next()?; + Ok(()) + } + + fn insert_counter(&self, counter: i64) -> Result<(), sqlite::Error> { + let conn = self.get_connection()?; + let mut statement = conn.prepare("INSERT INTO RATAT ( COUNTER ) VALUES ( :counter )")?; + statement.bind((":counter", counter))?; + statement.next()?; + Ok(()) + } + + fn select_counter(&self) -> Result<i64, sqlite::Error> { + let conn = self.get_connection()?; + let mut statement = conn.prepare("SELECT COUNTER FROM RATAT")?; + statement.next()?; + statement.read("COUNTER") + } + + fn update_counter(&self, counter: i64) -> Result<(), sqlite::Error> { + let conn = self.get_connection()?; + let mut statement = conn.prepare("UPDATE RATAT SET COUNTER = :counter")?; + statement.bind((":counter", counter))?; + statement.next()?; + Ok(()) + } +} + impl Default for SQLite { fn default() -> Self { SQLite { @@ -17,19 +60,47 @@ impl Default for SQLite { } impl Store for SQLite { - fn init(&mut self) -> Result<()> { + /// Init SQLite for counting. Creates `RATAT` table with single row, with + /// counter set to `0` + fn init(&mut self) -> Result<(), ContextError> { let conn_path = self.location.clone(); - self.connection = sqlite::open(conn_path).into_diagnostic().ok(); + let conn = sqlite::open(conn_path)?; + self.connection = Option::Some(conn); + + self.create_table()?; + self.insert_counter(0)?; + Ok(()) } + + /// Selects counter from `RATAT` table. Since sqlite doesn't scrutinize the + /// value and we have a more restricted datatype (i64 -> u8), we use + /// `try_from` which may produce an error result. + fn get_counter(&self) -> Result<u8, ContextError> { + let result = self.select_counter()?; + + let count = u8::try_from(result)?; + Ok(count) + } + + /// Updates row with counter value on `RATAT` table + fn set_counter(&mut self, new_count: u8) -> Result<u8, ContextError> { + self.update_counter(i64::from(new_count))?; + + Ok(new_count) + } } #[cfg(test)] mod tests { use super::*; + // Practically, the `Store` trait is what we interface with externally, and + // provides the more general `Result` interface we want to work with so + // that's what we test here + #[test] - fn sqlite_default() { + fn default() { let context = SQLite::default(); assert_eq!(context.location, ":memory:"); @@ -37,7 +108,7 @@ mod tests { } #[test] - fn sqlite_init() { + fn init() { let mut context = SQLite::default(); let res = context.init(); @@ -45,4 +116,84 @@ mod tests { assert!(res.is_ok()); assert!(context.connection.is_some()); } + + #[test] + fn get_counter() { + let mut context = SQLite::default(); + assert!(context.init().is_ok()); + + let res = context.get_counter(); + + assert!(res.is_ok()); + assert_eq!(res.unwrap(), 0); + } + + #[test] + fn get_counter_error() { + let context = SQLite::default(); + // don't initialize + + let res = context.get_counter(); + + assert!(res.is_err()); + } + + #[test] + fn get_counter_outside_range() { + let mut context = SQLite::default(); + assert!(context.init().is_ok()); + + let new_value = -1; + if let Some(conn) = &context.connection { + let mut statement = conn.prepare("UPDATE RATAT SET COUNTER = ?").unwrap(); + statement.bind((1, new_value)).unwrap(); + statement.next().unwrap(); + } + let res = context.get_counter(); + + assert!(res.is_err()); + } + + #[test] + fn set_counter() { + let mut context = SQLite::default(); + assert!(context.init().is_ok()); + + let expected = 240; + let res = context.set_counter(expected); + + assert!(res.is_ok()); + assert_eq!(res.unwrap(), expected); + } + + #[test] + fn set_counter_error() { + let mut context = SQLite::default(); + // don't initialize + + let res = context.set_counter(1); + + assert!(res.is_err()); + } + + #[test] + fn set_and_get() { + let mut context = SQLite::default(); + assert!(context.init().is_ok()); + + let count_0 = context.get_counter().unwrap(); + + assert_eq!(count_0, 0); + + let new_num = 244; + let res = context.set_counter(new_num); + + assert!(res.is_ok()); + assert_eq!(res.unwrap(), new_num); + + let count_1 = context.get_counter().unwrap(); + + assert_eq!(count_1, new_num); + assert_ne!(count_0, count_1); + } } |
