use std::{
collections::BTreeMap,
ffi::OsStr,
io::BufWriter,
path::{Path, PathBuf},
sync::{Arc, RwLock},
};
use crate::{
authentication_storage::{AuthenticationStorageError, StorageBackend},
Authentication,
};
#[derive(Clone, Debug)]
struct FileStorageCache {
content: BTreeMap<String, Authentication>,
}
#[derive(Clone, Debug)]
pub struct FileStorage {
pub path: PathBuf,
cache: Arc<RwLock<FileStorageCache>>,
}
#[derive(thiserror::Error, Debug)]
pub enum FileStorageError {
#[error(transparent)]
IOError(#[from] std::io::Error),
#[error("failed to parse {0}: {1}")]
JSONError(PathBuf, serde_json::Error),
}
impl FileStorageCache {
pub fn from_path(path: &Path) -> Result<Self, FileStorageError> {
match fs_err::read_to_string(path) {
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(Self {
content: BTreeMap::new(),
}),
Err(e) => Err(FileStorageError::IOError(e)),
Ok(content) => {
let content = serde_json::from_str(&content)
.map_err(|e| FileStorageError::JSONError(path.to_path_buf(), e))?;
Ok(Self { content })
}
}
}
}
impl FileStorage {
pub fn from_path(path: PathBuf) -> Result<Self, FileStorageError> {
let cache = Arc::new(RwLock::new(FileStorageCache::from_path(&path)?));
Ok(Self { path, cache })
}
#[cfg(feature = "dirs")]
pub fn new() -> Result<Self, FileStorageError> {
let home_dir = dirs::home_dir().ok_or_else(|| {
FileStorageError::IOError(std::io::Error::new(
std::io::ErrorKind::NotFound,
"Could not determine the home directory. Please ensure the $HOME environment variable is set.",
))
})?;
let path = home_dir.join(".rattler").join("credentials.json");
Self::from_path(path)
}
fn read_json(&self) -> Result<BTreeMap<String, Authentication>, FileStorageError> {
let new_cache = FileStorageCache::from_path(&self.path)?;
let mut cache = self.cache.write().unwrap();
cache.content = new_cache.content;
Ok(cache.content.clone())
}
fn write_json(&self, dict: &BTreeMap<String, Authentication>) -> Result<(), FileStorageError> {
let parent = self
.path
.parent()
.ok_or(FileStorageError::IOError(std::io::Error::new(
std::io::ErrorKind::NotFound,
"Parent directory not found",
)))?;
std::fs::create_dir_all(parent)?;
let prefix = self
.path
.file_stem()
.unwrap_or_else(|| OsStr::new("credentials"));
let extension = self
.path
.extension()
.and_then(OsStr::to_str)
.unwrap_or("json");
let mut temp_file = tempfile::Builder::new()
.prefix(prefix)
.suffix(&format!(".{extension}"))
.tempfile_in(parent)?;
serde_json::to_writer(BufWriter::new(&mut temp_file), dict)
.map_err(std::io::Error::from)?;
temp_file
.persist(&self.path)
.map_err(std::io::Error::from)?;
let mut cache = self.cache.write().unwrap();
cache.content = dict.clone();
Ok(())
}
}
impl StorageBackend for FileStorage {
fn store(
&self,
host: &str,
authentication: &crate::Authentication,
) -> Result<(), AuthenticationStorageError> {
let mut dict = self.read_json()?;
dict.insert(host.to_string(), authentication.clone());
Ok(self.write_json(&dict)?)
}
fn get(&self, host: &str) -> Result<Option<crate::Authentication>, AuthenticationStorageError> {
let cache = self.cache.read().unwrap();
Ok(cache.content.get(host).cloned())
}
fn delete(&self, host: &str) -> Result<(), AuthenticationStorageError> {
let mut dict = self.read_json()?;
if dict.remove(host).is_some() {
Ok(self.write_json(&dict)?)
} else {
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use std::{fs, io::Write};
use insta::assert_snapshot;
use tempfile::tempdir;
use super::*;
#[test]
fn test_file_storage() {
let file = tempdir().unwrap();
let path = file.path().join("test.json");
let storage = FileStorage::from_path(path.clone()).unwrap();
assert_eq!(storage.get("test").unwrap(), None);
storage
.store("test", &Authentication::CondaToken("password".to_string()))
.unwrap();
assert_eq!(
storage.get("test").unwrap(),
Some(Authentication::CondaToken("password".to_string()))
);
storage
.store(
"bearer",
&Authentication::BearerToken("password".to_string()),
)
.unwrap();
storage
.store(
"basic",
&Authentication::BasicHTTP {
username: "user".to_string(),
password: "password".to_string(),
},
)
.unwrap();
assert_snapshot!(fs::read_to_string(&path).unwrap());
storage.delete("test").unwrap();
assert_eq!(storage.get("test").unwrap(), None);
let mut file = std::fs::File::create(&path).unwrap();
file.write_all(b"invalid json").unwrap();
assert!(FileStorage::from_path(path.clone()).is_err());
}
}