use crate::error::WalletStorageError;
use aes_gcm::Aes256Gcm;
use log::*;
use std::{
fmt::{Display, Error, Formatter},
sync::Arc,
};
use tari_comms::{
multiaddr::Multiaddr,
peer_manager::NodeIdentity,
tor::TorIdentity,
types::{CommsPublicKey, CommsSecretKey},
};
const LOG_TARGET: &str = "wallet::database";
pub trait WalletBackend: Send + Sync + Clone {
fn fetch(&self, key: &DbKey) -> Result<Option<DbValue>, WalletStorageError>;
fn write(&self, op: WriteOperation) -> Result<Option<DbValue>, WalletStorageError>;
fn apply_encryption(&self, cipher: Aes256Gcm) -> Result<(), WalletStorageError>;
fn remove_encryption(&self) -> Result<(), WalletStorageError>;
}
#[derive(Debug, Clone, PartialEq)]
pub enum DbKey {
CommsSecretKey,
CommsPublicKey,
CommsAddress,
CommsFeatures,
Identity,
TorId,
ClientKey(String),
}
pub enum DbValue {
CommsSecretKey(CommsSecretKey),
CommsPublicKey(CommsPublicKey),
CommsAddress(Multiaddr),
CommsFeatures(u64),
Identity(NodeIdentity),
TorId(TorIdentity),
ClientValue(String),
ValueCleared,
}
#[derive(Clone)]
pub enum DbKeyValuePair {
CommsSecretKey(CommsSecretKey),
ClientKeyValue(String, String),
Identity(Box<NodeIdentity>),
TorId(TorIdentity),
}
pub enum WriteOperation {
Insert(DbKeyValuePair),
Remove(DbKey),
}
#[derive(Clone)]
pub struct WalletDatabase<T>
where T: WalletBackend + 'static
{
db: Arc<T>,
}
impl<T> WalletDatabase<T>
where T: WalletBackend + 'static
{
pub fn new(db: T) -> Self {
Self { db: Arc::new(db) }
}
pub async fn get_comms_secret_key(&self) -> Result<Option<CommsSecretKey>, WalletStorageError> {
let db_clone = self.db.clone();
let c = tokio::task::spawn_blocking(move || match db_clone.fetch(&DbKey::CommsSecretKey) {
Ok(None) => Ok(None),
Ok(Some(DbValue::CommsSecretKey(k))) => Ok(Some(k)),
Ok(Some(other)) => unexpected_result(DbKey::CommsSecretKey, other),
Err(e) => log_error(DbKey::CommsSecretKey, e),
})
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(c)
}
pub async fn set_comms_secret_key(&self, key: CommsSecretKey) -> Result<(), WalletStorageError> {
let db_clone = self.db.clone();
tokio::task::spawn_blocking(move || {
db_clone.write(WriteOperation::Insert(DbKeyValuePair::CommsSecretKey(key)))
})
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(())
}
pub async fn get_tor_id(&self) -> Result<Option<TorIdentity>, WalletStorageError> {
let db_clone = self.db.clone();
let c = tokio::task::spawn_blocking(move || match db_clone.fetch(&DbKey::TorId) {
Ok(None) => Ok(None),
Ok(Some(DbValue::TorId(k))) => Ok(Some(k)),
Ok(Some(other)) => unexpected_result(DbKey::TorId, other),
Err(e) => log_error(DbKey::CommsSecretKey, e),
})
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(c)
}
pub async fn set_tor_identity(&self, id: TorIdentity) -> Result<(), WalletStorageError> {
let db_clone = self.db.clone();
tokio::task::spawn_blocking(move || db_clone.write(WriteOperation::Insert(DbKeyValuePair::TorId(id))))
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(())
}
pub async fn clear_comms_secret_key(&self) -> Result<(), WalletStorageError> {
let db_clone = self.db.clone();
tokio::task::spawn_blocking(move || db_clone.write(WriteOperation::Remove(DbKey::CommsSecretKey)))
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(())
}
pub async fn apply_encryption(&self, cipher: Aes256Gcm) -> Result<(), WalletStorageError> {
let db_clone = self.db.clone();
tokio::task::spawn_blocking(move || db_clone.apply_encryption(cipher))
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))
.and_then(|inner_result| inner_result)
}
pub async fn remove_encryption(&self) -> Result<(), WalletStorageError> {
let db_clone = self.db.clone();
tokio::task::spawn_blocking(move || db_clone.remove_encryption())
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))
.and_then(|inner_result| inner_result)
}
pub async fn set_client_key_value(&self, key: String, value: String) -> Result<(), WalletStorageError> {
let db_clone = self.db.clone();
tokio::task::spawn_blocking(move || {
db_clone.write(WriteOperation::Insert(DbKeyValuePair::ClientKeyValue(key, value)))
})
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(())
}
pub async fn get_client_key_value(&self, key: String) -> Result<Option<String>, WalletStorageError> {
let db_clone = self.db.clone();
let c = tokio::task::spawn_blocking(move || match db_clone.fetch(&DbKey::ClientKey(key.clone())) {
Ok(None) => Ok(None),
Ok(Some(DbValue::ClientValue(k))) => Ok(Some(k)),
Ok(Some(other)) => unexpected_result(DbKey::ClientKey(key), other),
Err(e) => log_error(DbKey::ClientKey(key), e),
})
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(c)
}
pub async fn clear_client_value(&self, key: String) -> Result<bool, WalletStorageError> {
let db_clone = self.db.clone();
let c = tokio::task::spawn_blocking(move || {
match db_clone.write(WriteOperation::Remove(DbKey::ClientKey(key.clone()))) {
Ok(None) => Ok(false),
Ok(Some(DbValue::ValueCleared)) => Ok(true),
Ok(Some(other)) => unexpected_result(DbKey::ClientKey(key), other),
Err(e) => log_error(DbKey::ClientKey(key), e),
}
})
.await
.map_err(|err| WalletStorageError::BlockingTaskSpawnError(err.to_string()))??;
Ok(c)
}
}
impl Display for DbKey {
fn fmt(&self, f: &mut Formatter) -> Result<(), Error> {
match self {
DbKey::CommsSecretKey => f.write_str(&"CommsSecretKey".to_string()),
DbKey::CommsPublicKey => f.write_str(&"CommsPublicKey".to_string()),
DbKey::CommsAddress => f.write_str(&"CommsAddress".to_string()),
DbKey::CommsFeatures => f.write_str(&"Node features".to_string()),
DbKey::Identity => f.write_str(&"NodeIdentity".to_string()),
DbKey::TorId => f.write_str(&"TorId".to_string()),
DbKey::ClientKey(k) => f.write_str(&format!("ClientKey: {:?}", k)),
}
}
}
impl Display for DbValue {
fn fmt(&self, f: &mut Formatter) -> Result<(), Error> {
match self {
DbValue::CommsSecretKey(k) => f.write_str(&format!("CommsSecretKey: {:?}", k)),
DbValue::CommsPublicKey(k) => f.write_str(&format!("CommsPublicKey: {:?}", k)),
DbValue::ClientValue(v) => f.write_str(&format!("ClientValue: {:?}", v)),
DbValue::ValueCleared => f.write_str(&"ValueCleared".to_string()),
DbValue::CommsFeatures(_) => f.write_str(&"Node features".to_string()),
DbValue::CommsAddress(_) => f.write_str(&"Comms Address".to_string()),
DbValue::TorId(v) => f.write_str(&format!("Tor ID: {}", v)),
DbValue::Identity(v) => f.write_str(&format!("Node Identity: {}", v)),
}
}
}
fn log_error<T>(req: DbKey, err: WalletStorageError) -> Result<T, WalletStorageError> {
error!(
target: LOG_TARGET,
"Database access error on request: {}: {}",
req,
err.to_string()
);
Err(err)
}
fn unexpected_result<T>(req: DbKey, res: DbValue) -> Result<T, WalletStorageError> {
let msg = format!("Unexpected result for database query {}. Response: {}", req, res);
error!(target: LOG_TARGET, "{}", msg);
Err(WalletStorageError::UnexpectedResult(msg))
}
#[cfg(test)]
mod test {
use crate::storage::{
database::{WalletBackend, WalletDatabase},
memory_db::WalletMemoryDatabase,
sqlite_db::WalletSqliteDatabase,
sqlite_utilities::run_migration_and_create_sqlite_connection,
};
use rand::rngs::OsRng;
use tari_comms::types::CommsSecretKey;
use tari_crypto::keys::SecretKey;
use tari_test_utils::random::string;
use tempfile::tempdir;
use tokio::runtime::Runtime;
pub fn test_database_crud<T: WalletBackend + 'static>(backend: T) {
let mut runtime = Runtime::new().unwrap();
let db = WalletDatabase::new(backend);
assert!(runtime.block_on(db.get_comms_secret_key()).unwrap().is_none());
let secret_key = CommsSecretKey::random(&mut OsRng);
runtime.block_on(db.set_comms_secret_key(secret_key.clone())).unwrap();
let stored_key = runtime.block_on(db.get_comms_secret_key()).unwrap().unwrap();
assert_eq!(secret_key, stored_key);
runtime.block_on(db.clear_comms_secret_key()).unwrap();
assert!(runtime.block_on(db.get_comms_secret_key()).unwrap().is_none());
let client_key_values = vec![
("key1".to_string(), "value1".to_string()),
("key2".to_string(), "value2".to_string()),
("key3".to_string(), "value3".to_string()),
];
for kv in client_key_values.iter() {
runtime
.block_on(db.set_client_key_value(kv.0.clone(), kv.1.clone()))
.unwrap();
}
assert!(runtime
.block_on(db.get_client_key_value("wrong".to_string()))
.unwrap()
.is_none());
runtime
.block_on(db.set_client_key_value(client_key_values[0].0.clone(), "updated".to_string()))
.unwrap();
assert_eq!(
runtime
.block_on(db.get_client_key_value(client_key_values[0].0.clone()))
.unwrap()
.unwrap(),
"updated".to_string()
);
assert!(!runtime.block_on(db.clear_client_value("wrong".to_string())).unwrap());
assert!(runtime
.block_on(db.clear_client_value(client_key_values[0].0.clone()))
.unwrap());
assert!(!runtime
.block_on(db.clear_client_value(client_key_values[0].0.clone()))
.unwrap());
}
#[test]
fn test_database_crud_memory_db() {
test_database_crud(WalletMemoryDatabase::new());
}
#[test]
fn test_database_crud_sqlite_db() {
let db_name = format!("{}.sqlite3", string(8).as_str());
let db_folder = tempdir().unwrap().path().to_str().unwrap().to_string();
let connection = run_migration_and_create_sqlite_connection(&format!("{}{}", db_folder, db_name)).unwrap();
test_database_crud(WalletSqliteDatabase::new(connection, None).unwrap());
}
}