use std::{
fs::File,
io::{BufRead, BufReader, BufWriter, Write},
path::Path,
sync::Arc,
};
use parking_lot::{Mutex, RwLock};
use whasher::{GxPapayaMap, new_papaya_map};
use super::{
AclPassword, RespAclCategories, acl_exception::AclError, acl_parser::AclParser, user::User,
user_handle::UserHandle,
};
pub const DEFAULT_USER_NAME: &str = "default";
type UsersMap = GxPapayaMap<String, Arc<UserHandle>>;
pub struct AccessControlList {
users: RwLock<Arc<UsersMap>>,
default_user: RwLock<Option<Arc<UserHandle>>>,
save_lock: Mutex<()>,
}
impl AccessControlList {
pub fn new(
default_password: &str,
acl_configuration_file: Option<&Path>,
) -> Result<Self, AclError> {
let acl = Self::scratch();
if let Some(file) = acl_configuration_file {
acl.load(default_password, &file.display().to_string())?;
} else {
let handle = acl.create_default_user_handle(default_password)?;
*acl.default_user.write() = Some(handle);
}
Ok(acl)
}
fn scratch() -> Self {
Self {
users: RwLock::new(Arc::new(new_papaya_map())),
default_user: RwLock::new(None),
save_lock: Mutex::new(()),
}
}
pub fn get_user_handle(&self, username: &str) -> Option<Arc<UserHandle>> {
let users = Arc::clone(&self.users.read());
users.pin().get(username).map(Arc::clone)
}
pub fn get_default_user_handle(&self) -> Option<Arc<UserHandle>> {
self.default_user.read().clone()
}
pub fn add_user_handle(&self, user_handle: Arc<UserHandle>) -> Result<(), AclError> {
let username = user_handle.user().name.clone();
if Arc::clone(&self.users.read())
.pin()
.try_insert(username.clone(), user_handle)
.is_err()
{
return Err(AclError::UserAlreadyExists(username));
}
Ok(())
}
pub fn delete_user_handle(&self, username: &str) -> Result<bool, AclError> {
if username == DEFAULT_USER_NAME {
return Err(AclError::Acl(
"The special 'default' user cannot be removed from the system".into(),
));
}
Ok(
Arc::clone(&self.users.read())
.pin()
.remove(username)
.is_some(),
)
}
pub fn clear_users(&self) {
*self.users.write() = Arc::new(new_papaya_map());
}
pub fn get_user_handles(&self) -> Vec<(String, Arc<UserHandle>)> {
let users = Arc::clone(&self.users.read());
users
.pin()
.iter()
.map(|(name, handle)| (name.clone(), Arc::clone(handle)))
.collect()
}
pub fn len(&self) -> usize {
Arc::clone(&self.users.read()).pin().len()
}
pub fn is_empty(&self) -> bool {
Arc::clone(&self.users.read()).pin().is_empty()
}
pub fn create_default_user_handle(
&self,
default_password: &str,
) -> Result<Arc<UserHandle>, AclError> {
loop {
if let Some(handle) = self.get_user_handle(DEFAULT_USER_NAME) {
return Ok(handle);
}
let default_user = User::new(DEFAULT_USER_NAME.to_string());
default_user.add_category(RespAclCategories::ALL)?;
default_user.set_enabled(true);
if !default_password.is_empty() {
default_user.add_password_hash(AclPassword::from_string(default_password));
} else {
default_user.set_passwordless(true);
}
let default_user_handle = Arc::new(UserHandle::new(Arc::new(default_user)));
match self.add_user_handle(Arc::clone(&default_user_handle)) {
Ok(()) => return Ok(default_user_handle),
Err(AclError::UserAlreadyExists(_)) => continue,
Err(e) => return Err(e),
}
}
}
pub fn load(&self, default_password: &str, acl_configuration_file: &str) -> Result<(), AclError> {
if !Path::new(acl_configuration_file).exists() {
return Err(AclError::Acl(format!(
"Cannot find ACL configuration file '{acl_configuration_file}'"
)));
}
let scratch = Self::scratch();
let reader = File::open(acl_configuration_file)
.map(BufReader::new)
.map_err(|_| {
AclError::Acl(format!(
"Unable to open ACL configuration file '{acl_configuration_file}'"
))
})?;
if let Err(exception) = scratch.import(reader, acl_configuration_file) {
let AclError::Parsing {
message,
filename,
line,
} = exception
else {
return Err(exception);
};
return Err(AclError::parsing_wrap(&message, &filename, line));
}
let default_handle = scratch.create_default_user_handle(default_password)?;
*self.default_user.write() = Some(default_handle);
*self.users.write() = scratch.users.into_inner();
Ok(())
}
pub fn save(&self, acl_configuration_file: &str) -> Result<(), AclError> {
if acl_configuration_file.is_empty() {
return Err(AclError::Acl("ACL configuration file not set.".into()));
}
let _guard = self.save_lock.lock();
let file = File::create(acl_configuration_file).map_err(|e| AclError::Acl(e.to_string()))?;
let mut writer = BufWriter::with_capacity(1 << 16, file);
for (_, user_handle) in self.get_user_handles() {
writeln!(writer, "{}", user_handle.user().describe_user())
.map_err(|e| AclError::Acl(e.to_string()))?;
}
writer.flush().map_err(|e| AclError::Acl(e.to_string()))?;
Ok(())
}
fn import<R: BufRead>(&self, input: R, configuration_file: &str) -> Result<(), AclError> {
for (cur_line, line) in input.lines().enumerate() {
let line = line.map_err(|e| AclError::Acl(e.to_string()))?;
let line = line.trim();
if line.is_empty() || line.starts_with('#') {
continue;
}
if let Err(exception) = AclParser::parse_acl_rule(line, Some(self)) {
return Err(AclError::Parsing {
message: exception.to_string(),
filename: configuration_file.to_string(),
line: cur_line as i32 + 1,
});
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::{fs, thread};
use wresp::RespCommand;
use super::*;
#[test]
fn empty_input_file_yields_only_default() {
let dir = tempfile::tempdir().unwrap();
let file = dir.path().join("users.acl");
fs::write(&file, "").unwrap();
let acl = AccessControlList::new("", Some(&file)).unwrap();
let names: Vec<String> = acl.get_user_handles().into_iter().map(|(n, _)| n).collect();
assert_eq!(names, vec!["default".to_string()]);
}
#[test]
fn no_default_rule_creates_default() {
let dir = tempfile::tempdir().unwrap();
let file = dir.path().join("users.acl");
fs::write(
&file,
"user testA on >password123 +@admin\r\nuser testB on >passw0rd >password +@admin ",
)
.unwrap();
let acl = AccessControlList::new("", Some(&file)).unwrap();
assert_eq!(acl.len(), 3);
let names: Vec<String> = acl.get_user_handles().into_iter().map(|(n, _)| n).collect();
for expected in ["default", "testA", "testB"] {
assert!(names.iter().any(|n| n == expected), "missing {expected}");
}
let default = acl.get_default_user_handle().unwrap();
assert!(default.user().is_passwordless());
assert!(default.user().can_access_command(RespCommand::Get));
}
#[test]
fn with_default_rule_takes_precedence() {
let dir = tempfile::tempdir().unwrap();
let file = dir.path().join("users.acl");
fs::write(
&file,
"user testA on >password123 +@admin +@slow\r\nuser testB on >passw0rd >password +@admin\r\nuser default on nopass +@admin +@slow",
)
.unwrap();
let acl = AccessControlList::new("ignored-password", Some(&file)).unwrap();
assert_eq!(acl.len(), 3);
let described = acl
.get_default_user_handle()
.unwrap()
.user()
.describe_user();
assert!(
!described.contains('#'),
"no password expected: {described}"
);
assert!(described.contains("nopass"));
}
#[test]
fn load_replaces_users_atomically() {
let dir = tempfile::tempdir().unwrap();
let file = dir.path().join("users.acl");
fs::write(
&file,
"user testA on >password123 +@admin +@slow\r\nuser testB on >passw0rd >password +@admin +@slow\r\nuser testC on >passw0rd\r\nuser default on nopass +@admin +@slow",
)
.unwrap();
let acl = AccessControlList::new("", Some(&file)).unwrap();
assert_eq!(acl.len(), 4);
fs::write(
&file,
"user testD on >password123\r\nuser testB on >passw0rd +@admin +@slow\r\nuser default on nopass +@admin",
)
.unwrap();
acl.load("", &file.display().to_string()).unwrap();
assert_eq!(acl.len(), 3);
assert!(acl.get_user_handle("testA").is_none());
assert!(acl.get_user_handle("testC").is_none());
assert!(acl.get_user_handle("testD").is_some());
assert!(acl.get_user_handle("testB").is_some());
assert!(acl.get_default_user_handle().is_some());
}
#[test]
fn load_errors() {
let acl = AccessControlList::new("", None).unwrap();
let err = acl.load("", "/nonexistent/users.acl").unwrap_err();
assert!(
err
.to_string()
.contains("Cannot find ACL configuration file")
);
let dir = tempfile::tempdir().unwrap();
let file = dir.path().join("users.acl");
fs::write(&file, "# 注释行\nuser alice on +@nosuch\n").unwrap();
let err = match AccessControlList::new("", Some(&file)) {
Err(e) => e,
Ok(_) => panic!("expected parse failure"),
};
let msg = err.to_string();
assert!(
msg.starts_with("Unable to parse ACL rule") && msg.contains(":2:"),
"{msg}"
);
}
#[test]
fn save_load_roundtrip() {
let dir = tempfile::tempdir().unwrap();
let file = dir.path().join("users.acl");
let acl = AccessControlList::new("passw0rd", None).unwrap();
AclParser::parse_acl_rule("user alice on >secret +@keyspace +set", Some(&acl)).unwrap();
acl.save(&file.display().to_string()).unwrap();
let restored = AccessControlList::new("", Some(&file)).unwrap();
assert_eq!(restored.len(), 2);
let alice = restored.get_user_handle("alice").unwrap().user();
assert!(alice.is_enabled());
assert!(alice.validate_password(&AclPassword::from_string("secret")));
assert!(alice.can_access_command(RespCommand::Set));
assert!(
restored
.get_default_user_handle()
.unwrap()
.user()
.validate_password(&AclPassword::from_string("passw0rd"))
);
let err = acl.save("").unwrap_err();
assert!(err.to_string().contains("ACL configuration file not set"));
}
#[test]
fn user_handle_crud() {
let acl = AccessControlList::new("", None).unwrap();
let handle = Arc::new(UserHandle::new(Arc::new(User::new("bob".into()))));
acl.add_user_handle(Arc::clone(&handle)).unwrap();
let dup = Arc::new(UserHandle::new(Arc::new(User::new("bob".into()))));
assert!(matches!(
acl.add_user_handle(dup),
Err(AclError::UserAlreadyExists(u)) if u == "bob"
));
assert!(acl.delete_user_handle("bob").unwrap());
assert!(!acl.delete_user_handle("bob").unwrap());
let err = acl.delete_user_handle("default").unwrap_err();
assert!(err.to_string().contains("cannot be removed"));
acl.clear_users();
assert!(acl.is_empty());
}
#[test]
fn create_default_user_handle_converges() {
let acl = Arc::new(AccessControlList::new("", None).unwrap());
let handles: Vec<_> = (0..8)
.map(|_| {
let acl = Arc::clone(&acl);
thread::spawn(move || acl.create_default_user_handle("").unwrap())
})
.collect();
let joined: Vec<_> = handles.into_iter().map(|h| h.join().unwrap()).collect();
let first = &joined[0];
for h in &joined {
assert!(Arc::ptr_eq(first, h));
}
}
}