use std::cell::{Cell, RefCell};
use super::Server;
use super::pubsub::{self, Kind};
pub(crate) mod class {
pub(crate) const KEYSPACE: u32 = 1 << 0;
pub(crate) const KEYEVENT: u32 = 1 << 1;
pub(crate) const GENERIC: u32 = 1 << 2;
pub(crate) const STRING: u32 = 1 << 3;
pub(crate) const LIST: u32 = 1 << 4;
pub(crate) const SET: u32 = 1 << 5;
pub(crate) const HASH: u32 = 1 << 6;
pub(crate) const ZSET: u32 = 1 << 7;
pub(crate) const EXPIRED: u32 = 1 << 8;
pub(crate) const EVICTED: u32 = 1 << 9;
pub(crate) const STREAM: u32 = 1 << 10;
pub(crate) const KEY_MISS: u32 = 1 << 11;
pub(crate) const MODULE: u32 = 1 << 13;
pub(crate) const NEW: u32 = 1 << 14;
pub(crate) const OVERWRITTEN: u32 = 1 << 15;
pub(crate) const TYPE_CHANGED: u32 = 1 << 16;
pub(crate) const SUBKEYSPACE: u32 = 1 << 19;
pub(crate) const SUBKEYEVENT: u32 = 1 << 20;
pub(crate) const SUBKEYSPACEITEM: u32 = 1 << 21;
pub(crate) const SUBKEYSPACEEVENT: u32 = 1 << 22;
pub(crate) const ARRAY: u32 = 1 << 23;
pub(crate) const ALL: u32 =
GENERIC | STRING | LIST | SET | HASH | ZSET | EXPIRED | EVICTED | STREAM | MODULE | ARRAY;
}
const CHANNELS: u32 = class::KEYSPACE | class::KEYEVENT;
const FIRST: &[(u8, u32)] = &[
(b'g', class::GENERIC),
(b'$', class::STRING),
(b'l', class::LIST),
(b's', class::SET),
(b'h', class::HASH),
(b'z', class::ZSET),
(b'x', class::EXPIRED),
(b'e', class::EVICTED),
(b't', class::STREAM),
(b'd', class::MODULE),
(b'a', class::ARRAY),
(b'n', class::NEW),
(b'o', class::OVERWRITTEN),
(b'c', class::TYPE_CHANGED),
];
const SECOND: &[(u8, u32)] = &[
(b'K', class::KEYSPACE),
(b'E', class::KEYEVENT),
(b'm', class::KEY_MISS),
(b'S', class::SUBKEYSPACE),
(b'T', class::SUBKEYEVENT),
(b'I', class::SUBKEYSPACEITEM),
(b'V', class::SUBKEYSPACEEVENT),
];
pub(crate) const ACCEPTED: &str = "Ag$lshzxeKEtmdnocaSTIV";
const WIDEST: usize = FIRST.len() + SECOND.len();
pub(crate) fn parse(classes: &[u8]) -> Option<u32> {
let mut flags = 0;
for &c in classes {
flags |= match c {
b'A' => class::ALL,
_ => {
let found = FIRST
.iter()
.chain(SECOND)
.find(|(ch, _)| *ch == c)
.map(|(_, bit)| *bit);
found?
}
};
}
Some(flags)
}
pub(crate) fn format(flags: u32) -> ([u8; WIDEST], usize) {
let mut out = [0u8; WIDEST];
let mut len = 0;
if flags & class::ALL == class::ALL {
out[len] = b'A';
len += 1;
} else {
for (c, bit) in FIRST {
if flags & bit != 0 {
out[len] = *c;
len += 1;
}
}
}
for (c, bit) in SECOND {
if flags & bit != 0 {
out[len] = *c;
len += 1;
}
}
(out, len)
}
struct Event {
db: usize,
name: &'static str,
key: Vec<u8>,
}
thread_local! {
static PENDING: RefCell<Vec<Event>> = const { RefCell::new(Vec::new()) };
static ARMED: Cell<u32> = const { Cell::new(0) };
}
pub(super) fn arm(server: &Server) -> u32 {
let flags = server.notify_flags();
let live = flags & CHANNELS != 0 && server.anyone_subscribed();
ARMED.replace(if live { flags } else { 0 })
}
pub(crate) fn armed() -> bool {
ARMED.get() != 0
}
pub(crate) fn fire(db: usize, class: u32, name: &'static str, key: &[u8]) {
if ARMED.get() & class == 0 {
return;
}
keep(db, name, key);
}
#[cold]
#[inline(never)]
fn keep(db: usize, name: &'static str, key: &[u8]) {
let event = yo_alloc::allow(|| Event {
db,
name,
key: key.to_vec(),
});
PENDING.with_borrow_mut(|pending| yo_alloc::allow(|| pending.push(event)));
}
pub(super) fn drain(server: &Server, was: u32) {
let flags = ARMED.replace(was);
if flags == 0 {
return;
}
let events = PENDING.with_borrow_mut(std::mem::take);
if events.is_empty() {
return;
}
for event in &events {
send(server, flags, event);
}
PENDING.with_borrow_mut(|pending| {
if pending.is_empty() {
let mut events = events;
events.clear();
*pending = events;
}
});
}
fn send(server: &Server, flags: u32, event: &Event) {
let mut channel = [0u8; CHANNEL_MAX];
if flags & class::KEYSPACE != 0 {
let head = prefix(&mut channel, b"__keyspace@", event.db);
yo_alloc::allow(|| {
let mut name = channel[..head].to_vec();
name.extend_from_slice(&event.key);
pubsub::deliver(server, Kind::Channel, &name, event.name.as_bytes());
});
}
if flags & class::KEYEVENT != 0 {
let head = prefix(&mut channel, b"__keyevent@", event.db);
yo_alloc::allow(|| {
let mut name = channel[..head].to_vec();
name.extend_from_slice(event.name.as_bytes());
pubsub::deliver(server, Kind::Channel, &name, &event.key);
});
}
}
const CHANNEL_MAX: usize = 11 + 20 + 3;
fn prefix(into: &mut [u8; CHANNEL_MAX], head: &[u8], db: usize) -> usize {
into[..head.len()].copy_from_slice(head);
let mut at = head.len();
let mut digits = [0u8; 20];
let mut n = db;
let mut count = 0;
loop {
digits[count] = b'0' + (n % 10) as u8;
count += 1;
n /= 10;
if n == 0 {
break;
}
}
for i in (0..count).rev() {
into[at] = digits[i];
at += 1;
}
into[at..at + 3].copy_from_slice(b"__:");
at + 3
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_accepted_characters_are_the_ones_the_error_message_lists() {
for c in ACCEPTED.bytes() {
assert!(parse(&[c]).is_some(), "{} was refused", c as char);
}
for c in 0u8..=127 {
if !ACCEPTED.as_bytes().contains(&c) {
assert!(parse(&[c]).is_none(), "{} was accepted", c as char);
}
}
assert_eq!(parse(b""), Some(0), "the empty setting is off and is legal");
}
#[test]
fn a_setting_reads_back_in_one_order_whatever_order_it_was_set_in() {
let round = |s: &str| {
let flags = parse(s.as_bytes()).unwrap();
let (buf, len) = format(flags);
String::from_utf8(buf[..len].to_vec()).unwrap()
};
assert_eq!(round("KEA"), "AKE");
assert_eq!(round("AKE"), "AKE");
assert_eq!(round("gxE"), "gxE");
assert_eq!(round("K"), "K");
assert_eq!(round("E"), "E");
assert_eq!(round(""), "");
assert_eq!(round("A"), "A");
assert_eq!(round("Kg"), "gK");
assert_eq!(round("nKE"), "nKE");
assert_eq!(round("KEg$lshzxetdmn"), "g$lshzxetdnKEm");
}
#[test]
fn all_leaves_out_the_four_classes_it_leaves_out() {
let all = parse(b"A").unwrap();
for (c, bit) in [
(b'm', class::KEY_MISS),
(b'n', class::NEW),
(b'o', class::OVERWRITTEN),
(b'c', class::TYPE_CHANGED),
] {
assert_eq!(all & bit, 0, "A should not contain {}", c as char);
}
let both = parse(b"An").unwrap();
assert_eq!(both & class::NEW, class::NEW);
assert_eq!(both & class::EXPIRED, class::EXPIRED);
let (buf, len) = format(both);
assert_eq!(&buf[..len], b"A");
let (buf, len) = format(parse(b"Am").unwrap());
assert_eq!(&buf[..len], b"Am");
}
#[test]
fn a_channel_name_carries_the_database_it_happened_on() {
let mut buf = [0u8; CHANNEL_MAX];
let len = prefix(&mut buf, b"__keyspace@", 0);
assert_eq!(&buf[..len], b"__keyspace@0__:");
let len = prefix(&mut buf, b"__keyevent@", 15);
assert_eq!(&buf[..len], b"__keyevent@15__:");
}
}