use std::sync::Mutex;
use std::sync::atomic::Ordering::{Acquire, Relaxed, Release};
use std::sync::atomic::{AtomicBool, AtomicU64};
use yo_common::{Code, Error, Result, glob_matches, sha256};
use super::args::{self, Args, is};
use super::keyspec::{self, Access};
use super::table::{self, Spec};
use super::{Server, Session};
use crate::reply::Out;
pub(super) const DEFAULT: &[u8] = b"default";
const CATEGORIES: [&str; 31] = [
"keyspace",
"read",
"write",
"set",
"sortedset",
"list",
"hash",
"string",
"array",
"bitmap",
"hyperloglog",
"geo",
"stream",
"pubsub",
"admin",
"fast",
"slow",
"blocking",
"dangerous",
"connection",
"transaction",
"scripting",
"bloom",
"cms",
"cuckoo",
"graph",
"json",
"search",
"tdigest",
"timeseries",
"topk",
];
const WORDS: usize = table::count().div_ceil(64);
const READ: u8 = 1;
const WRITE: u8 = 2;
#[derive(Debug, Clone, PartialEq, Eq)]
struct Pattern {
flags: u8,
glob: Vec<u8>,
}
impl Pattern {
fn describe(&self, into: &mut Vec<u8>) {
match self.flags {
READ => into.extend_from_slice(b"%R~"),
WRITE => into.extend_from_slice(b"%W~"),
_ => into.push(b'~'),
}
into.extend_from_slice(&self.glob);
}
}
#[derive(Debug, Clone)]
struct Selector {
allowed: [u64; WORDS],
all_commands: bool,
future: bool,
all_keys: bool,
all_channels: bool,
rules: Vec<u8>,
firstargs: Vec<(u16, Vec<Vec<u8>>)>,
deniedfirst: Vec<(u16, Vec<Vec<u8>>)>,
patterns: Vec<Pattern>,
channels: Vec<Vec<u8>>,
}
impl Selector {
fn new() -> Selector {
Selector {
allowed: [0; WORDS],
all_commands: false,
future: false,
all_keys: false,
all_channels: false,
rules: Vec::new(),
firstargs: Vec::new(),
deniedfirst: Vec::new(),
patterns: Vec::new(),
channels: Vec::new(),
}
}
fn bit(&self, i: usize) -> bool {
i < table::count() && self.allowed[i / 64] & (1 << (i % 64)) != 0
}
fn set(&mut self, i: usize, allow: bool) {
if i >= table::count() {
return;
}
if allow {
self.allowed[i / 64] |= 1 << (i % 64);
} else {
self.allowed[i / 64] &= !(1 << (i % 64));
self.all_commands = false;
}
self.firstargs.retain(|(at, _)| usize::from(*at) != i);
self.deniedfirst.retain(|(at, _)| usize::from(*at) != i);
}
fn note(&mut self, rule: &[u8], allow: bool) {
let mut kept: Vec<u8> = Vec::with_capacity(self.rules.len() + rule.len() + 2);
for old in self.rules.split(|b| *b == b' ') {
if old.is_empty() {
continue;
}
let name = &old[1..];
let same = name == rule;
let child =
name.len() > rule.len() && name.starts_with(rule) && name[rule.len()] == b'|';
if same || child {
continue;
}
if !kept.is_empty() {
kept.push(b' ');
}
kept.extend_from_slice(old);
}
if !kept.is_empty() {
kept.push(b' ');
}
kept.push(if allow { b'+' } else { b'-' });
kept.extend_from_slice(rule);
self.rules = kept;
}
fn allow_first(&mut self, i: u16, first: &[u8]) {
let lower = first.to_ascii_lowercase();
match self.firstargs.binary_search_by_key(&i, |(at, _)| *at) {
Ok(at) => {
let list = &mut self.firstargs[at].1;
if !list.contains(&lower) {
list.push(lower);
}
}
Err(at) => self.firstargs.insert(at, (i, vec![lower])),
}
}
fn firsts(&self, i: u16) -> &[Vec<u8>] {
match self.firstargs.binary_search_by_key(&i, |(at, _)| *at) {
Ok(at) => &self.firstargs[at].1,
Err(_) => &[],
}
}
fn deny_first(&mut self, i: u16, first: &[u8]) {
let lower = first.to_ascii_lowercase();
self.all_commands = false;
match self.deniedfirst.binary_search_by_key(&i, |(at, _)| *at) {
Ok(at) => {
let list = &mut self.deniedfirst[at].1;
if !list.contains(&lower) {
list.push(lower);
}
}
Err(at) => self.deniedfirst.insert(at, (i, vec![lower])),
}
}
fn denied(&self, i: u16) -> &[Vec<u8>] {
match self.deniedfirst.binary_search_by_key(&i, |(at, _)| *at) {
Ok(at) => &self.deniedfirst[at].1,
Err(_) => &[],
}
}
fn reset_commands(&mut self, all: bool) {
self.allowed = [if all { u64::MAX } else { 0 }; WORDS];
self.all_commands = all;
self.future = all;
self.rules.clear();
self.firstargs.clear();
self.deniedfirst.clear();
}
fn describe_commands(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(self.rules.len() + 8);
out.extend_from_slice(if self.future { b"+@all" } else { b"-@all" });
if !self.rules.is_empty() {
out.push(b' ');
out.extend_from_slice(&self.rules);
}
out
}
fn describe_keys(&self) -> Vec<u8> {
let mut out = Vec::new();
if self.all_keys {
out.extend_from_slice(b"~*");
return out;
}
for pattern in &self.patterns {
if !out.is_empty() {
out.push(b' ');
}
pattern.describe(&mut out);
}
out
}
fn describe_channels(&self) -> Vec<u8> {
let mut out = Vec::new();
if self.all_channels {
out.extend_from_slice(b"&*");
return out;
}
for channel in &self.channels {
if !out.is_empty() {
out.push(b' ');
}
out.push(b'&');
out.extend_from_slice(channel);
}
out
}
fn describe(&self) -> Vec<u8> {
let mut out = self.describe_keys();
if !out.is_empty() {
out.push(b' ');
}
if self.all_channels {
out.extend_from_slice(b"&* ");
} else {
out.extend_from_slice(b"resetchannels ");
for channel in &self.channels {
out.push(b'&');
out.extend_from_slice(channel);
out.push(b' ');
}
}
out.extend_from_slice(&self.describe_commands());
out
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Bad {
Syntax,
Unknown,
AfterAllKeys,
AfterAllChannels,
NoSuchPassword,
Hash,
NestedFirstArg,
Unmatched,
}
impl Bad {
fn text(self) -> &'static str {
match self {
Bad::Syntax | Bad::Unmatched => "Syntax error",
Bad::Unknown => "Unknown command or category name in ACL",
Bad::AfterAllKeys => {
"Adding a pattern after the * pattern (or the 'allkeys' flag) is not valid and does not have any effect. Try 'resetkeys' to start with an empty list of patterns"
}
Bad::AfterAllChannels => {
"Adding a pattern after the * pattern (or the 'allchannels' flag) is not valid and does not have any effect. Try 'resetchannels' to start with an empty list of channels"
}
Bad::NoSuchPassword => {
"The password you are trying to remove from the user does not exist"
}
Bad::Hash => {
"The password hash must be exactly 64 characters and contain only lowercase hexadecimal characters"
}
Bad::NestedFirstArg => "Allowing first-arg of a subcommand is not supported",
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct User {
name: Vec<u8>,
enabled: bool,
nopass: bool,
skip_sanitize: bool,
passwords: Vec<[u8; 64]>,
selectors: Vec<Selector>,
}
impl User {
fn new(name: &[u8]) -> User {
User {
name: name.to_vec(),
enabled: false,
nopass: false,
skip_sanitize: false,
passwords: Vec::new(),
selectors: vec![Selector::new()],
}
}
fn default_user() -> User {
let mut user = User::new(DEFAULT);
user.enabled = true;
user.nopass = true;
let root = &mut user.selectors[0];
root.reset_commands(true);
root.all_keys = true;
root.all_channels = true;
user
}
fn flags(&self) -> Vec<&'static str> {
let mut out = Vec::with_capacity(3);
out.push(if self.enabled { "on" } else { "off" });
if self.nopass {
out.push("nopass");
}
out.push(if self.skip_sanitize {
"skip-sanitize-payload"
} else {
"sanitize-payload"
});
out
}
fn describe(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(64);
out.extend_from_slice(&self.name);
for flag in self.flags() {
out.push(b' ');
out.extend_from_slice(flag.as_bytes());
}
for hash in &self.passwords {
out.extend_from_slice(b" #");
out.extend_from_slice(hash);
}
for (at, selector) in self.selectors.iter().enumerate() {
out.push(b' ');
if at == 0 {
out.extend_from_slice(&selector.describe());
} else {
out.push(b'(');
out.extend_from_slice(&selector.describe());
out.push(b')');
}
}
out
}
fn unrestricted(&self) -> bool {
self.selectors.len() == 1
&& self.selectors[0].all_commands
&& self.selectors[0].all_keys
&& self.selectors[0].all_channels
}
fn admits(&self, password: &[u8]) -> bool {
if !self.enabled {
return false;
}
if self.nopass {
return true;
}
let guess = sha256::hex(password);
let mut hit = false;
for hash in &self.passwords {
hit |= same(hash, &guess);
}
hit
}
}
fn same(a: &[u8], b: &[u8]) -> bool {
let mut diff = u8::from(a.len() != b.len());
for i in 0..a.len().max(b.len()) {
diff |= a.get(i).copied().unwrap_or(0) ^ b.get(i).copied().unwrap_or(0xff);
}
diff == 0
}
#[derive(Debug)]
pub(crate) struct Users {
table: Mutex<Vec<User>>,
generation: AtomicU64,
guarded: AtomicBool,
restricted: AtomicBool,
}
impl Default for Users {
fn default() -> Users {
Users {
table: Mutex::new(vec![User::default_user()]),
generation: AtomicU64::new(1),
guarded: AtomicBool::new(false),
restricted: AtomicBool::new(false),
}
}
}
impl Users {
fn with<T>(&self, f: impl FnOnce(&mut Vec<User>) -> (bool, T)) -> T {
let mut held = self.table.lock().unwrap_or_else(|e| e.into_inner());
let (changed, out) = f(&mut held);
if changed {
let guarded = held
.binary_search_by(|u| u.name.as_slice().cmp(DEFAULT))
.is_ok_and(|at| !held[at].nopass || !held[at].enabled);
self.guarded.store(guarded, Relaxed);
self.restricted
.store(held.iter().any(|u| !u.unrestricted()), Relaxed);
self.generation.fetch_add(1, Release);
}
out
}
fn get(&self, name: &[u8]) -> Option<User> {
self.with(|table| {
let found = table
.binary_search_by(|u| u.name.as_slice().cmp(name))
.ok()
.map(|at| table[at].clone());
(false, found)
})
}
fn generation(&self) -> u64 {
self.generation.load(Acquire)
}
}
impl Server {
#[must_use]
pub(crate) fn guarded(&self) -> bool {
self.acl.guarded.load(Relaxed)
}
#[must_use]
pub(crate) fn restricted(&self) -> bool {
self.acl.restricted.load(Relaxed)
}
pub fn set_password(&self, password: &[u8]) {
self.acl.with(|table| {
let Ok(at) = table.binary_search_by(|u| u.name.as_slice().cmp(DEFAULT)) else {
return (false, ());
};
let user = &mut table[at];
user.passwords.clear();
if password.is_empty() {
user.nopass = true;
} else {
user.nopass = false;
user.passwords.push(sha256::hex(password));
}
(true, ())
});
self.plain.set(password);
}
pub(crate) fn with_password<T>(&self, f: impl FnOnce(&[u8]) -> T) -> T {
self.plain.with(f)
}
pub(crate) fn users(&self) -> &Users {
&self.acl
}
}
#[derive(Debug, Default)]
pub(crate) struct Plain {
secret: Mutex<Vec<u8>>,
}
impl Plain {
fn set(&self, password: &[u8]) {
let mut held = self.secret.lock().unwrap_or_else(|e| e.into_inner());
held.clear();
held.extend_from_slice(password);
}
fn with<T>(&self, f: impl FnOnce(&[u8]) -> T) -> T {
let held = self.secret.lock().unwrap_or_else(|e| e.into_inner());
f(&held)
}
}
pub(super) fn authenticate(
server: &Server,
session: &mut Session,
user: &[u8],
password: &[u8],
) -> bool {
let stamp = server.acl.generation();
let Some(found) = server.acl.get(user) else {
return false;
};
if !found.admits(password) {
return false;
}
session.become_user(stamp, found);
session.admit(true);
true
}
fn set_user(user: &mut User, rule: &[u8]) -> std::result::Result<(), Bad> {
if rule.is_empty() {
return Ok(());
}
if word(rule, b"on") {
user.enabled = true;
} else if word(rule, b"off") {
user.enabled = false;
} else if word(rule, b"skip-sanitize-payload") {
user.skip_sanitize = true;
} else if word(rule, b"sanitize-payload") {
user.skip_sanitize = false;
} else if word(rule, b"nopass") {
user.nopass = true;
user.passwords.clear();
} else if word(rule, b"resetpass") {
user.nopass = false;
user.passwords.clear();
} else if rule[0] == b'>' || rule[0] == b'#' {
let hash = hash_of(rule)?;
if !user.passwords.contains(&hash) {
user.passwords.push(hash);
}
user.nopass = false;
} else if rule[0] == b'<' || rule[0] == b'!' {
let hash = hash_of(rule)?;
let before = user.passwords.len();
user.passwords.retain(|held| *held != hash);
if user.passwords.len() == before {
return Err(Bad::NoSuchPassword);
}
} else if rule[0] == b'(' && rule[rule.len() - 1] == b')' {
let mut selector = Selector::new();
for word in split(&rule[1..rule.len() - 1]) {
set_selector(&mut selector, &word)?;
}
user.selectors.push(selector);
} else if rule[0] == b'(' {
return Err(Bad::Unmatched);
} else if word(rule, b"clearselectors") {
user.selectors.truncate(1);
} else if word(rule, b"reset") {
let name = std::mem::take(&mut user.name);
*user = User::new(&name);
} else {
return set_selector(&mut user.selectors[0], rule);
}
Ok(())
}
fn hash_of(rule: &[u8]) -> std::result::Result<[u8; 64], Bad> {
let rest = &rule[1..];
if rule[0] == b'>' || rule[0] == b'<' {
return Ok(sha256::hex(rest));
}
let ok = rest.len() == 64
&& rest
.iter()
.all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(b));
if !ok {
return Err(Bad::Hash);
}
let mut hash = [0u8; 64];
hash.copy_from_slice(rest);
Ok(hash)
}
fn set_selector(selector: &mut Selector, rule: &[u8]) -> std::result::Result<(), Bad> {
if word(rule, b"allkeys") || rule == b"~*" {
selector.all_keys = true;
selector.patterns.clear();
} else if word(rule, b"resetkeys") {
selector.all_keys = false;
selector.patterns.clear();
} else if word(rule, b"allchannels") || rule == b"&*" {
selector.all_channels = true;
selector.channels.clear();
} else if word(rule, b"resetchannels") {
selector.all_channels = false;
selector.channels.clear();
} else if word(rule, b"allcommands") || rule == b"+@all" {
selector.reset_commands(true);
} else if word(rule, b"nocommands") || rule == b"-@all" {
selector.reset_commands(false);
} else if rule[0] == b'~' || rule[0] == b'%' {
add_pattern(selector, rule)?;
} else if rule[0] == b'&' {
if selector.all_channels {
return Err(Bad::AfterAllChannels);
}
let glob = &rule[1..];
if glob.contains(&b' ') {
return Err(Bad::Syntax);
}
if !selector.channels.iter().any(|held| held == glob) {
selector.channels.push(glob.to_vec());
}
} else if rule[0] == b'+' && rule.get(1) != Some(&b'@') {
add_command(selector, &rule[1..], true)?;
} else if rule[0] == b'-' && rule.get(1) != Some(&b'@') {
add_command(selector, &rule[1..], false)?;
} else if rule[0] == b'+' || rule[0] == b'-' {
let allow = rule[0] == b'+';
add_category(selector, &rule[2..], allow)?;
let mut name = Vec::with_capacity(rule.len() - 1);
name.push(b'@');
name.extend_from_slice(&rule[2..].to_ascii_lowercase());
selector.note(&name, allow);
} else {
return Err(Bad::Syntax);
}
Ok(())
}
fn add_pattern(selector: &mut Selector, rule: &[u8]) -> std::result::Result<(), Bad> {
if selector.all_keys {
return Err(Bad::AfterAllKeys);
}
let mut flags = 0u8;
let mut at = 1;
if rule[0] == b'%' {
let mut ok = true;
while at < rule.len() {
let letter = rule[at].to_ascii_uppercase();
if letter == b'R' && flags & READ == 0 {
flags |= READ;
} else if letter == b'W' && flags & WRITE == 0 {
flags |= WRITE;
} else if rule[at] == b'~' {
at += 1;
break;
} else {
ok = false;
break;
}
at += 1;
}
if flags == 0 || !ok {
return Err(Bad::Syntax);
}
} else {
flags = READ | WRITE;
}
let glob = &rule[at..];
if glob.contains(&b' ') {
return Err(Bad::Syntax);
}
if let Some(held) = selector.patterns.iter_mut().find(|p| p.glob == glob) {
held.flags |= flags;
} else {
selector.patterns.push(Pattern {
flags,
glob: glob.to_vec(),
});
}
Ok(())
}
fn add_command(selector: &mut Selector, name: &[u8], allow: bool) -> std::result::Result<(), Bad> {
let Some(bar) = name.iter().rposition(|b| *b == b'|') else {
let Some(spec) = table::lookup(name) else {
return Err(Bad::Unknown);
};
selector.set(table::index_of(spec), allow);
selector.note(&name.to_ascii_lowercase(), allow);
return Ok(());
};
let (head, first) = (&name[..bar], &name[bar + 1..]);
let Some(spec) = table::lookup(head) else {
return Err(Bad::Unknown);
};
if head.contains(&b'|') {
return Err(Bad::NestedFirstArg);
}
if first.is_empty() {
return Err(Bad::Syntax);
}
let at = u16::try_from(table::index_of(spec)).unwrap_or(u16::MAX);
if allow {
if !selector.bit(usize::from(at)) {
selector.allow_first(at, first);
}
} else {
if !super::CONTAINERS.contains(&spec.name) {
return Err(Bad::Unknown);
}
selector.deny_first(at, first);
}
selector.note(&name.to_ascii_lowercase(), allow);
Ok(())
}
fn add_category(selector: &mut Selector, name: &[u8], allow: bool) -> std::result::Result<(), Bad> {
let Some(wanted) = category(name) else {
return Err(Bad::Unknown);
};
for (at, spec) in table::COMMANDS.iter().enumerate() {
if spec.acl.iter().any(|held| &held[1..] == wanted) {
selector.set(at, allow);
}
}
for (container, sub, cats) in table::SUBCATS {
if !cats.iter().any(|held| &held[1..] == wanted) {
continue;
}
let Some(spec) = table::lookup(container.as_bytes()) else {
continue;
};
let at = u16::try_from(table::index_of(spec)).unwrap_or(u16::MAX);
if !allow {
selector.deny_first(at, sub.as_bytes());
} else if !selector.bit(usize::from(at)) {
selector.allow_first(at, sub.as_bytes());
}
}
Ok(())
}
fn category(name: &[u8]) -> Option<&'static str> {
CATEGORIES
.iter()
.find(|held| name.eq_ignore_ascii_case(held.as_bytes()))
.copied()
}
fn word(rule: &[u8], keyword: &[u8]) -> bool {
rule.eq_ignore_ascii_case(keyword)
}
fn split(inside: &[u8]) -> Vec<Vec<u8>> {
inside
.split(|b| *b == b' ')
.filter(|part| !part.is_empty())
.map(<[u8]>::to_vec)
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub(crate) enum Denied {
Command,
Key(usize),
Channel(usize),
}
impl Denied {
fn rank(self) -> u8 {
match self {
Denied::Command => 1,
Denied::Key(_) => 2,
Denied::Channel(_) => 4,
}
}
fn at(self) -> usize {
match self {
Denied::Command => 0,
Denied::Key(at) | Denied::Channel(at) => at,
}
}
}
struct Channels {
first: usize,
count: Option<usize>,
pattern: bool,
}
fn channels_of(name: &str) -> Option<Channels> {
let spec = match name {
"subscribe" | "ssubscribe" => Channels {
first: 1,
count: None,
pattern: false,
},
"psubscribe" => Channels {
first: 1,
count: None,
pattern: true,
},
"publish" | "spublish" => Channels {
first: 1,
count: Some(1),
pattern: false,
},
_ => return None,
};
Some(spec)
}
fn key_ok(selector: &Selector, key: &[u8], need: Access) -> bool {
if selector.all_keys {
return true;
}
selector
.patterns
.iter()
.any(|p| Access::from_bits(p.flags).covers(need) && glob_matches(&p.glob, key))
}
fn channel_ok(selector: &Selector, channel: &[u8], pattern: bool) -> bool {
if selector.all_channels {
return true;
}
selector.channels.iter().any(|held| {
if pattern {
held == channel
} else {
glob_matches(held, channel)
}
})
}
fn selector_ok(
selector: &Selector,
spec: &'static Spec,
args: Args<'_>,
base: usize,
) -> std::result::Result<(), Denied> {
if !selector.all_commands && !spec.flags.contains(&"no_auth") {
let at = table::index_of(spec);
let index = u16::try_from(at).unwrap_or(u16::MAX);
let sub = (args.len() > base + 1).then(|| args.get(base + 1));
let named = |list: &[Vec<u8>]| {
sub.is_some_and(|word| list.iter().any(|held| word.eq_ignore_ascii_case(held)))
};
if selector.bit(at) {
if named(selector.denied(index)) {
return Err(Denied::Command);
}
} else if !named(selector.firsts(index)) {
return Err(Denied::Command);
}
}
if !selector.all_keys && keyspec::takes_keys(spec, args, base) {
let mut refused = None;
keyspec::find(spec, args, base, &mut |run| {
if refused.is_some() {
return;
}
let need = run.need();
for i in 0..run.count {
let at = run.first + i * run.step;
if at < args.len() && !key_ok(selector, args.get(at), need) {
refused = Some(Denied::Key(at));
return;
}
}
});
if let Some(why) = refused {
return Err(why);
}
}
if !selector.all_channels
&& let Some(where_) = channels_of(spec.name)
{
let first = base + where_.first;
let stop = where_
.count
.map_or(args.len(), |n| (first + n).min(args.len()));
for at in first..stop {
if !channel_ok(selector, args.get(at), where_.pattern) {
return Err(Denied::Channel(at));
}
}
}
Ok(())
}
pub(crate) fn permits(
user: &User,
spec: &'static Spec,
args: Args<'_>,
base: usize,
) -> std::result::Result<(), Denied> {
if let Some(root) = user.selectors.first()
&& root.all_commands
&& root.all_keys
&& root.all_channels
{
return Ok(());
}
let mut worst = Denied::Command;
for selector in &user.selectors {
match selector_ok(selector, spec, args, base) {
Ok(()) => return Ok(()),
Err(why) => {
if why.rank() > worst.rank()
|| (why.rank() == worst.rank() && why.at() > worst.at())
{
worst = why;
}
}
}
}
Err(worst)
}
pub(crate) fn refusal(
why: Denied,
user: &[u8],
spec: &'static Spec,
args: Args<'_>,
base: usize,
verbose: bool,
) -> String {
let name = String::from_utf8_lossy(user);
match why {
Denied::Command => {
let sub = container_name(spec, args, base);
format!("User {name} has no permissions to run the '{sub}' command")
}
Denied::Key(at) if verbose => {
let key = String::from_utf8_lossy(args.get(at));
format!("User {name} has no permissions to access the '{key}' key")
}
Denied::Key(_) => "No permissions to access a key".to_string(),
Denied::Channel(at) if verbose => {
let channel = String::from_utf8_lossy(args.get(at));
format!("User {name} has no permissions to access the '{channel}' channel")
}
Denied::Channel(_) => "No permissions to access a channel".to_string(),
}
}
fn container_name(spec: &'static Spec, args: Args<'_>, base: usize) -> String {
if super::CONTAINERS.contains(&spec.name) && args.len() > base + 1 {
let sub = String::from_utf8_lossy(args.get(base + 1)).to_lowercase();
return format!("{}|{sub}", spec.name);
}
spec.name.to_string()
}
#[derive(Debug)]
pub(crate) struct Identity {
name: Vec<u8>,
stamp: u64,
user: User,
}
impl Default for Identity {
fn default() -> Identity {
Identity {
name: DEFAULT.to_vec(),
stamp: 0,
user: User::default_user(),
}
}
}
impl Session {
pub(crate) fn acl_name(&self) -> &[u8] {
&self.acl.name
}
pub(super) fn become_user(&mut self, stamp: u64, user: User) {
yo_alloc::allow(|| {
self.acl.name.clear();
self.acl.name.extend_from_slice(&user.name);
});
self.acl.stamp = stamp;
self.acl.user = user;
self.sock.set_text(|text| &mut text.user, &self.acl.name);
}
pub(super) fn forget_user(&mut self) {
*self.acl = Identity::default();
self.sock.set_text(|text| &mut text.user, b"");
}
fn acl_stamp(&self) -> u64 {
self.acl.stamp
}
fn acl_refresh(&mut self, now: u64, fresh: Option<User>) {
self.acl.stamp = now;
if let Some(user) = fresh {
self.acl.user = user;
}
}
fn acl_cached(&self) -> &User {
&self.acl.user
}
}
fn current<'a>(server: &Server, session: &'a mut Session) -> &'a User {
let now = server.acl.generation();
if session.acl_stamp() != now {
let fresh = server.acl.get(session.acl_name());
session.acl_refresh(now, fresh);
}
session.acl_cached()
}
pub(super) fn gate(
server: &Server,
session: &mut Session,
spec: &'static Spec,
args: Args<'_>,
) -> Option<String> {
let user = current(server, session);
let why = permits(user, spec, args, 0).err()?;
let said = refusal(why, &user.name, spec, args, 0, false);
Some(yo_alloc::allow(|| format!("NOPERM {said}")))
}
const ARITIES: [(&[u8], i32); 9] = [
(b"cat", -2),
(b"deluser", -3),
(b"dryrun", -4),
(b"genpass", -2),
(b"getuser", 3),
(b"help", 2),
(b"list", 2),
(b"setuser", -3),
(b"users", 2),
];
pub(super) fn execute(
server: &Server,
session: &mut Session,
args: Args<'_>,
out: &mut Out,
) -> Result<()> {
let sub = args.get(1);
if let Some((name, arity)) = ARITIES
.iter()
.find(|(name, _)| sub.eq_ignore_ascii_case(name))
{
let n = args.len() as i32;
if (*arity > 0 && n != *arity) || (*arity < 0 && n < -*arity) {
let name = std::str::from_utf8(name).unwrap_or("acl");
return Err(args::wrong_arity_sub("acl", name));
}
}
if is(sub, b"whoami") {
if args.len() != 2 {
return Err(args::wrong_arity_sub("acl", "whoami"));
}
out.bulk(session.acl_name());
} else if is(sub, b"cat") {
cat(args, out)?;
} else if is(sub, b"list") {
yo_alloc::allow(|| {
server.users().with(|table| {
out.array(table.len());
for user in table.iter() {
let mut line = b"user ".to_vec();
line.extend_from_slice(&user.describe());
out.bulk(&line);
}
(false, ())
});
});
} else if is(sub, b"users") {
server.users().with(|table| {
out.array(table.len());
for user in table.iter() {
out.bulk(&user.name);
}
(false, ())
});
} else if is(sub, b"getuser") {
yo_alloc::allow(|| getuser(server, args, out));
} else if is(sub, b"setuser") {
return yo_alloc::allow(|| setuser(server, args, out));
} else if is(sub, b"deluser") {
return deluser(server, args, out);
} else if is(sub, b"genpass") {
genpass(args, out)?;
} else if is(sub, b"dryrun") {
return yo_alloc::allow(|| dryrun(server, args, out));
} else if is(sub, b"help") {
super::server::help(out, HELP);
} else {
return Err(args::unknown_subcommand(sub, "ACL"));
}
Ok(())
}
fn cat(args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() == 2 {
out.array(CATEGORIES.len());
for name in CATEGORIES {
out.bulk(name.as_bytes());
}
return Ok(());
}
if args.len() > 3 {
return Err(args::subcommand_syntax(args.get(1), "ACL"));
}
let Some(wanted) = category(args.get(2)) else {
return Err(yo_alloc::allow(|| {
Error::fmt(
Code::Invalid,
format_args!(
"Unknown category '{}'",
String::from_utf8_lossy(args.get(2))
),
)
}));
};
let start = out.len();
let mut n = 0;
for spec in table::COMMANDS {
if spec.acl.iter().any(|held| &held[1..] == wanted) {
out.bulk(spec.name.as_bytes());
n += 1;
}
}
out.close_array(start, n);
Ok(())
}
fn getuser(server: &Server, args: Args<'_>, out: &mut Out) {
let Some(user) = server.users().get(args.get(2)) else {
out.nil();
return;
};
out.map(6);
out.bulk(b"flags");
let flags = user.flags();
let root = &user.selectors[0];
out.set(flags.len());
for flag in &flags {
out.bulk(flag.as_bytes());
}
out.bulk(b"passwords");
out.array(user.passwords.len());
for hash in &user.passwords {
out.bulk(hash);
}
describe_selector(root, out);
out.bulk(b"selectors");
out.array(user.selectors.len() - 1);
for selector in &user.selectors[1..] {
out.map(3);
describe_selector(selector, out);
}
}
fn describe_selector(selector: &Selector, out: &mut Out) {
out.bulk(b"commands");
out.bulk(&selector.describe_commands());
out.bulk(b"keys");
out.bulk(&selector.describe_keys());
out.bulk(b"channels");
out.bulk(&selector.describe_channels());
}
fn setuser(server: &Server, args: Args<'_>, out: &mut Out) -> Result<()> {
let name = args.get(2);
if name.contains(&b' ') || name.contains(&0) {
return Err(Error::new(
Code::Invalid,
"Usernames can't contain spaces or null characters",
));
}
let rules = merge(args, 3)?;
let outcome = server.users().with(|table| {
let at = table.binary_search_by(|u| u.name.as_slice().cmp(name));
let mut staged = match at {
Ok(at) => table[at].clone(),
Err(_) => User::new(name),
};
for rule in &rules {
if let Err(why) = set_user(&mut staged, rule) {
return (false, Err((rule.clone(), why)));
}
}
match at {
Ok(at) => table[at] = staged,
Err(at) => table.insert(at, staged),
}
(true, Ok(()))
});
match outcome {
Ok(()) => {
out.ok();
Ok(())
}
Err((rule, why)) => Err(Error::fmt(
Code::Invalid,
format_args!(
"Error in ACL SETUSER modifier '{}': {}",
String::from_utf8_lossy(&rule),
why.text()
),
)),
}
}
fn merge(args: Args<'_>, from: usize) -> Result<Vec<Vec<u8>>> {
let mut out: Vec<Vec<u8>> = Vec::with_capacity(args.len() - from);
let mut open: Option<usize> = None;
for i in from..args.len() {
let word = args.get(i);
if open.is_none() && word.first() == Some(&b'(') && word.last() != Some(&b')') {
open = Some(i);
out.push(word.to_vec());
continue;
}
if open.is_some() {
let held = out.last_mut().expect("an open bracket left a rule behind");
held.push(b' ');
held.extend_from_slice(word);
if word.last() == Some(&b')') {
open = None;
}
continue;
}
out.push(word.to_vec());
}
if let Some(at) = open {
return Err(Error::fmt(
Code::Invalid,
format_args!(
"Unmatched parenthesis in acl selector starting at '{}'.",
String::from_utf8_lossy(args.get(at))
),
));
}
Ok(out)
}
fn deluser(server: &Server, args: Args<'_>, out: &mut Out) -> Result<()> {
for i in 2..args.len() {
if args.get(i) == DEFAULT {
return Err(Error::new(
Code::Invalid,
"The 'default' user cannot be removed",
));
}
}
let gone = server.users().with(|table| {
let mut gone = 0;
for i in 2..args.len() {
if let Ok(at) = table.binary_search_by(|u| u.name.as_slice().cmp(args.get(i))) {
table.remove(at);
gone += 1;
}
}
(gone > 0, gone)
});
out.int(gone);
Ok(())
}
fn genpass(args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() > 3 {
return Err(args::subcommand_syntax(args.get(1), "ACL"));
}
let bits = if args.len() == 3 { args.int(2)? } else { 256 };
if bits <= 0 || bits > 4096 {
return Err(Error::new(
Code::Invalid,
"ACL GENPASS argument must be the number of bits for the output password, a positive number up to 4096",
));
}
let chars = ((bits + 3) / 4) as usize;
let mut raw = [0u8; 4096 / 8 + 1];
let bytes = chars.div_ceil(2);
yo_common::entropy::fill(&mut raw[..bytes]);
const DIGITS: &[u8; 16] = b"0123456789abcdef";
let mut hex = [0u8; 1024];
for (i, slot) in hex[..chars].iter_mut().enumerate() {
let byte = raw[i / 2];
*slot = DIGITS[usize::from(if i % 2 == 0 { byte >> 4 } else { byte & 0xf })];
}
out.bulk(&hex[..chars]);
Ok(())
}
fn dryrun(server: &Server, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(user) = server.users().get(args.get(2)) else {
return Err(Error::fmt(
Code::Invalid,
format_args!("User '{}' not found", String::from_utf8_lossy(args.get(2))),
));
};
let Some(spec) = table::lookup(args.get(3)) else {
return Err(Error::fmt(
Code::Invalid,
format_args!(
"Command '{}' not found",
String::from_utf8_lossy(args.get(3))
),
));
};
if !table::arity_ok(spec, args.len() - 3) {
return Err(args::wrong_arity(spec.name));
}
match permits(&user, spec, args, 3) {
Ok(()) => out.ok(),
Err(why) => {
let text = refusal(why, &user.name, spec, args, 3, true);
out.bulk(text.as_bytes());
}
}
Ok(())
}
const HELP: &[&str] = &[
"ACL <subcommand> [<arg> [value] [opt] ...]. Subcommands are:",
"CAT [<category>]",
" List all commands that belong to <category>, or all command categories",
" when no category is specified.",
"DELUSER <username> [<username> ...]",
" Delete a list of users.",
"DRYRUN <username> <command> [<arg> ...]",
" Returns whether the user can execute the given command without executing the command.",
"GETUSER <username>",
" Get the user's details.",
"GENPASS [<bits>]",
" Generate a secure 256-bit user password. The optional `bits` argument can",
" be used to specify a different size.",
"LIST",
" Show users details in config file format.",
"SETUSER <username> <attribute> [<attribute> ...]",
" Create or modify a user with the specified attributes.",
"USERS",
" List all the registered usernames.",
"WHOAMI",
" Return the current connection username.",
"HELP",
" Print this help.",
];
#[cfg(test)]
mod tests {
use super::*;
use crate::proto::Limits;
use crate::request::{Argv, Step};
fn built(rules: &[&str]) -> std::result::Result<User, Bad> {
let mut user = User::new(b"u");
for rule in rules {
set_user(&mut user, rule.as_bytes())?;
}
Ok(user)
}
fn listed(rules: &[&str]) -> String {
let user = built(rules).expect("every rule here is a good one");
String::from_utf8(user.describe()).expect("the description is text")
}
fn allows(rules: &[&str], words: &[&str]) -> bool {
let user = built(rules).expect("every rule here is a good one");
let mut buf = format!("*{}\r\n", words.len()).into_bytes();
for word in words {
buf.extend_from_slice(format!("${}\r\n{word}\r\n", word.len()).as_bytes());
}
let mut argv = Argv::new();
let Ok(Step::Command { .. }) = argv.decode(&buf, &Limits::default()) else {
panic!("the test wrote a command that does not decode");
};
let args = Args::new(&argv, &buf);
let spec = table::lookup(args.name()).expect("a command this server has");
permits(&user, spec, args, 0).is_ok()
}
#[test]
fn a_key_permission_with_no_pattern_after_it_is_an_empty_pattern() {
assert_eq!(
listed(&["%R"]),
"u off sanitize-payload %R~ resetchannels -@all"
);
assert_eq!(
listed(&["%RW"]),
"u off sanitize-payload ~ resetchannels -@all"
);
}
#[test]
fn a_key_permission_that_is_not_r_or_w_once_each_is_a_syntax_error() {
assert_eq!(built(&["%"]).err(), Some(Bad::Syntax));
assert_eq!(built(&["%~k:*"]).err(), Some(Bad::Syntax));
assert_eq!(built(&["%RR~k:*"]).err(), Some(Bad::Syntax));
assert_eq!(built(&["%X~k:*"]).err(), Some(Bad::Syntax));
}
#[test]
fn a_subcommand_can_be_taken_back_off_a_container_that_was_allowed() {
assert_eq!(
listed(&["+config", "-config|get"]),
"u off sanitize-payload resetchannels -@all +config -config|get"
);
assert!(allows(
&["+config", "-config|get"],
&["config", "set", "maxmemory", "0"]
));
assert!(!allows(
&["+config", "-config|get"],
&["config", "get", "maxmemory"]
));
assert_eq!(
listed(&["+config", "-config|get", "+config"]),
"u off sanitize-payload resetchannels -@all +config"
);
assert!(allows(
&["+config", "-config|get", "+config"],
&["config", "get", "maxmemory"]
));
}
#[test]
fn a_first_argument_can_only_be_taken_off_a_command_that_has_subcommands() {
assert_eq!(built(&["-get|nope"]).err(), Some(Bad::Unknown));
assert_eq!(built(&["-select|0"]).err(), Some(Bad::Unknown));
assert!(built(&["+get|nope"]).is_ok());
}
#[test]
fn a_category_a_subcommand_holds_reaches_that_subcommand_and_no_further() {
let deny = ["~*", "+@all", "-@admin"];
assert!(!allows(&deny, &["config", "get", "maxmemory"]));
assert!(allows(&deny, &["config", "help"]));
assert!(!allows(&deny, &["client", "kill", "id", "4"]));
assert!(allows(&deny, &["client", "setname", "x"]));
assert!(allows(&deny, &["acl", "whoami"]));
assert!(!allows(&deny, &["acl", "setuser", "u"]));
let grant = ["~*", "-@all", "+@admin"];
assert!(allows(&grant, &["config", "get", "maxmemory"]));
assert!(!allows(&grant, &["config", "help"]));
assert!(!allows(&grant, &["get", "k"]));
let config = table::lookup(b"config").expect("a command this server has");
assert_eq!(config.acl, ["@slow"]);
}
#[test]
fn the_flags_a_user_reports_are_the_ones_about_the_account() {
let user = built(&["on", "~*", "&*", "+@all"]).expect("good rules");
assert_eq!(user.flags(), ["on", "sanitize-payload"]);
}
}