use std::sync::{Arc, RwLock};
use serde::{Deserialize, Serialize};
use crate::bot::BotCredential;
use crate::error::AuthError;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct Credentials {
pub bot: Option<BotCredential>,
pub token: Option<String>,
}
pub trait CredentialStore: Send + Sync {
fn load(&self) -> Result<Option<Credentials>, AuthError>;
fn save(&self, credentials: &Credentials) -> Result<(), AuthError>;
fn clear(&self) -> Result<(), AuthError>;
}
#[derive(Debug, Clone, Default)]
pub struct MemoryCredentialStore {
inner: Arc<RwLock<Option<Credentials>>>,
}
impl MemoryCredentialStore {
pub fn new(initial: Credentials) -> Self {
Self {
inner: Arc::new(RwLock::new(
Some(initial).filter(|c| c.bot.is_some() || c.token.is_some()),
)),
}
}
pub fn shared(self) -> Arc<Self> {
Arc::new(self)
}
}
impl CredentialStore for MemoryCredentialStore {
fn load(&self) -> Result<Option<Credentials>, AuthError> {
Ok(self.inner.read().unwrap_or_else(|e| e.into_inner()).clone())
}
fn save(&self, credentials: &Credentials) -> Result<(), AuthError> {
let value = Some(credentials.clone()).filter(|c| c.bot.is_some() || c.token.is_some());
*self.inner.write().unwrap_or_else(|e| e.into_inner()) = value;
Ok(())
}
fn clear(&self) -> Result<(), AuthError> {
*self.inner.write().unwrap_or_else(|e| e.into_inner()) = None;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn bot(id: &str) -> BotCredential {
BotCredential::new(id.to_string(), "secret".into())
}
#[test]
fn memory_store_roundtrip() {
let store = MemoryCredentialStore::new(Credentials::default());
assert!(store.load().unwrap().is_none());
let creds = Credentials {
bot: Some(bot("bot1")),
token: Some("tok-1".into()),
};
store.save(&creds).unwrap();
let loaded = store.load().unwrap().unwrap();
assert_eq!(loaded.bot.as_ref().map(|b| b.id.as_str()), Some("bot1"));
assert_eq!(loaded.token.as_deref(), Some("tok-1"));
}
#[test]
fn memory_store_save_empty_clears() {
let store = MemoryCredentialStore::new(Credentials {
bot: Some(bot("bot1")),
token: None,
});
store.save(&Credentials::default()).unwrap();
assert!(store.load().unwrap().is_none());
}
#[test]
fn memory_store_clear() {
let store = MemoryCredentialStore::new(Credentials {
bot: Some(bot("bot1")),
token: Some("tok-1".into()),
});
store.clear().unwrap();
assert!(store.load().unwrap().is_none());
store.clear().unwrap();
}
}