use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering::Relaxed;
use super::table::Spec;
use super::{Args, Flow, Server, Session, args};
use crate::reply::Out;
use yo_common::lock::Lock;
use yo_common::{Error, Result, glob};
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum Kind {
Channel,
Pattern,
Shard,
}
impl Kind {
const fn joined(self) -> &'static [u8] {
match self {
Kind::Channel => b"subscribe",
Kind::Pattern => b"psubscribe",
Kind::Shard => b"ssubscribe",
}
}
const fn left(self) -> &'static [u8] {
match self {
Kind::Channel => b"unsubscribe",
Kind::Pattern => b"punsubscribe",
Kind::Shard => b"sunsubscribe",
}
}
const fn word(self) -> &'static [u8] {
match self {
Kind::Channel => b"message",
Kind::Pattern => b"pmessage",
Kind::Shard => b"smessage",
}
}
}
const KINDS: [Kind; 3] = [Kind::Channel, Kind::Pattern, Kind::Shard];
#[derive(Default)]
pub(crate) struct Subs {
channels: Vec<Vec<u8>>,
patterns: Vec<Vec<u8>>,
shard: Vec<Vec<u8>>,
}
impl Subs {
fn list(&self, kind: Kind) -> &Vec<Vec<u8>> {
match kind {
Kind::Channel => &self.channels,
Kind::Pattern => &self.patterns,
Kind::Shard => &self.shard,
}
}
fn list_mut(&mut self, kind: Kind) -> &mut Vec<Vec<u8>> {
match kind {
Kind::Channel => &mut self.channels,
Kind::Pattern => &mut self.patterns,
Kind::Shard => &mut self.shard,
}
}
fn total(&self) -> usize {
self.channels.len() + self.patterns.len() + self.shard.len()
}
fn reported(&self, kind: Kind) -> usize {
match kind {
Kind::Shard => self.shard.len(),
_ => self.channels.len() + self.patterns.len(),
}
}
}
impl Session {
pub(crate) fn subscribed(&self) -> bool {
self.subs.as_ref().is_some_and(|s| s.total() != 0)
}
fn sub_total(&self) -> usize {
self.subs.as_ref().map_or(0, |s| s.total())
}
fn sub_count(&self, kind: Kind) -> usize {
self.subs.as_ref().map_or(0, |s| s.reported(kind))
}
fn subs_mut(&mut self) -> &mut Subs {
if self.subs.is_none() {
self.subs = Some(yo_alloc::allow(Box::<Subs>::default));
}
self.subs.as_mut().unwrap()
}
}
#[derive(Clone, Copy)]
struct Row {
client: u64,
conn: u32,
thread: usize,
}
#[derive(Default)]
pub(crate) struct Registry {
channels: HashMap<Vec<u8>, Vec<Row>>,
patterns: HashMap<Vec<u8>, Vec<Row>>,
shard: HashMap<Vec<u8>, Vec<Row>>,
rows: usize,
clients: usize,
}
impl Registry {
fn table(&self, kind: Kind) -> &HashMap<Vec<u8>, Vec<Row>> {
match kind {
Kind::Channel => &self.channels,
Kind::Pattern => &self.patterns,
Kind::Shard => &self.shard,
}
}
fn table_mut(&mut self, kind: Kind) -> &mut HashMap<Vec<u8>, Vec<Row>> {
match kind {
Kind::Channel => &mut self.channels,
Kind::Pattern => &mut self.patterns,
Kind::Shard => &mut self.shard,
}
}
fn add(&mut self, kind: Kind, name: &[u8], row: Row) {
yo_alloc::allow(|| {
self.table_mut(kind)
.entry(name.to_vec())
.or_default()
.push(row)
});
self.rows += 1;
}
fn remove(&mut self, kind: Kind, name: &[u8], client: u64) {
let table = self.table_mut(kind);
let Some(rows) = table.get_mut(name) else {
return;
};
let mut gone = false;
if let Some(at) = rows.iter().position(|r| r.client == client) {
rows.swap_remove(at);
gone = true;
}
if rows.is_empty() {
table.remove(name);
}
if gone {
self.rows -= 1;
}
}
}
struct Body {
channel: Vec<u8>,
payload: Vec<u8>,
}
pub(crate) struct Envelope {
conn: u32,
client: u64,
kind: Kind,
pattern: Vec<u8>,
body: Arc<Body>,
}
impl Envelope {
pub(crate) const fn conn(&self) -> u32 {
self.conn
}
pub(crate) const fn client(&self) -> u64 {
self.client
}
pub(crate) fn write(&self, out: &mut Out) {
match self.kind {
Kind::Pattern => {
out.push(4);
out.bulk(Kind::Pattern.word());
out.bulk(&self.pattern);
}
kind => {
out.push(3);
out.bulk(kind.word());
}
}
out.bulk(&self.body.channel);
out.bulk(&self.body.payload);
}
}
#[derive(Default)]
#[repr(align(64))]
pub(crate) struct Mailbox {
queue: Lock<Vec<Envelope>>,
len: AtomicUsize,
here: AtomicUsize,
}
pub(crate) fn boxes(threads: usize) -> Box<[Mailbox]> {
(0..threads.max(1)).map(|_| Mailbox::default()).collect()
}
pub(crate) struct Counts {
pub(crate) clients: usize,
pub(crate) channels: usize,
pub(crate) patterns: usize,
pub(crate) shard: usize,
}
impl Server {
pub(crate) fn anyone_subscribed(&self) -> bool {
self.subs.load(Relaxed) != 0
}
fn note_subs(&self, reg: &Registry) {
self.subs.store(reg.rows, Relaxed);
}
fn note_here(&self, thread: usize, by: isize) {
let Some(mail) = self.mail.get(thread) else {
return;
};
let was = mail.here.load(Relaxed);
let now = if by < 0 {
was.saturating_sub(1)
} else {
was.saturating_add(1)
};
mail.here.store(now, Relaxed);
}
fn post(&self, thread: usize, env: Envelope) {
let Some(mail) = self.mail.get(thread) else {
return;
};
let mut queue = mail.queue.lock();
yo_alloc::allow(|| queue.push(env));
mail.len.store(queue.len(), Relaxed);
}
pub(crate) fn mail_here(&self) -> usize {
self.mail[self.my_slot()].len.load(Relaxed)
}
pub(crate) fn posted(&self) -> usize {
let mail = &self.mail[self.my_slot()];
mail.len.load(Relaxed) + mail.here.load(Relaxed)
}
pub(crate) fn take_mail(&self, into: &mut Vec<Envelope>) {
let at = self.my_slot();
let mut queue = self.mail[at].queue.lock();
yo_alloc::allow(|| into.append(&mut queue));
self.mail[at].len.store(0, Relaxed);
}
pub(crate) fn pubsub_counts(&self) -> Counts {
let reg = self.pubsub.lock();
Counts {
clients: reg.clients,
channels: reg.channels.len(),
patterns: reg.patterns.len(),
shard: reg.shard.len(),
}
}
}
pub(crate) fn execute(
server: &Server,
session: &mut Session,
spec: &'static Spec,
args: Args<'_>,
out: &mut Out,
) -> Result<Flow> {
match spec.name {
"subscribe" => join(server, session, args, out, Kind::Channel),
"psubscribe" => join(server, session, args, out, Kind::Pattern),
"ssubscribe" => join(server, session, args, out, Kind::Shard),
"unsubscribe" => leave(server, session, args, out, Kind::Channel),
"punsubscribe" => leave(server, session, args, out, Kind::Pattern),
"sunsubscribe" => leave(server, session, args, out, Kind::Shard),
"publish" => publish(server, args, out, Kind::Channel),
"spublish" => publish(server, args, out, Kind::Shard),
_ => introspect(server, args, out)?,
}
Ok(Flow::Continue)
}
fn join(server: &Server, session: &mut Session, args: Args<'_>, out: &mut Out, kind: Kind) {
let client = session.id();
let conn = session.conn;
let thread = server.my_slot();
let was = session.sub_total();
let mut reg = server.pubsub.lock();
for i in 1..args.len() {
let name = args.get(i);
let subs = session.subs_mut();
if !subs.list(kind).iter().any(|held| held == name) {
yo_alloc::allow(|| subs.list_mut(kind).push(name.to_vec()));
reg.add(
kind,
name,
Row {
client,
conn,
thread,
},
);
}
out.push(3);
out.bulk(kind.joined());
out.bulk(name);
out.uint(subs.reported(kind) as u64);
}
if was == 0 && session.sub_total() != 0 {
reg.clients += 1;
server.note_here(thread, 1);
}
server.note_subs(®);
}
fn leave(server: &Server, session: &mut Session, args: Args<'_>, out: &mut Out, kind: Kind) {
let client = session.id();
let thread = server.my_slot();
let was = session.sub_total();
let mut reg = server.pubsub.lock();
if args.len() > 1 {
for i in 1..args.len() {
let name = args.get(i);
drop_one(session, &mut reg, kind, client, name);
out.push(3);
out.bulk(kind.left());
out.bulk(name);
out.uint(session.sub_count(kind) as u64);
}
} else {
let held = session
.subs
.as_ref()
.map_or_else(Vec::new, |s| yo_alloc::allow(|| s.list(kind).clone()));
if held.is_empty() {
out.push(3);
out.bulk(kind.left());
out.nil();
out.uint(session.sub_count(kind) as u64);
}
for name in held {
drop_one(session, &mut reg, kind, client, &name);
out.push(3);
out.bulk(kind.left());
out.bulk(&name);
out.uint(session.sub_count(kind) as u64);
}
}
if was != 0 && session.sub_total() == 0 {
reg.clients -= 1;
server.note_here(thread, -1);
}
server.note_subs(®);
}
fn drop_one(session: &mut Session, reg: &mut Registry, kind: Kind, client: u64, name: &[u8]) {
let Some(subs) = session.subs.as_mut() else {
return;
};
let Some(at) = subs.list(kind).iter().position(|held| held == name) else {
return;
};
subs.list_mut(kind).swap_remove(at);
reg.remove(kind, name, client);
}
fn publish(server: &Server, args: Args<'_>, out: &mut Out, kind: Kind) {
if !server.anyone_subscribed() {
out.uint(0);
return;
}
out.uint(deliver(server, kind, args.get(1), args.get(2)));
}
pub(crate) fn deliver(server: &Server, kind: Kind, channel: &[u8], payload: &[u8]) -> u64 {
let reg = server.pubsub.lock();
let mut body: Option<Arc<Body>> = None;
let mut sent = 0u64;
if let Some(rows) = reg.table(kind).get(channel) {
for row in rows {
let body = shared(&mut body, channel, payload);
server.post(
row.thread,
Envelope {
conn: row.conn,
client: row.client,
kind,
pattern: Vec::new(),
body: Arc::clone(body),
},
);
sent += 1;
}
}
if kind == Kind::Channel {
for (pattern, rows) in ®.patterns {
if !glob::matches(pattern, channel) {
continue;
}
for row in rows {
let body = shared(&mut body, channel, payload);
let env = yo_alloc::allow(|| Envelope {
conn: row.conn,
client: row.client,
kind: Kind::Pattern,
pattern: pattern.clone(),
body: Arc::clone(body),
});
server.post(row.thread, env);
sent += 1;
}
}
}
sent
}
fn shared<'a>(body: &'a mut Option<Arc<Body>>, channel: &[u8], payload: &[u8]) -> &'a Arc<Body> {
body.get_or_insert_with(|| {
yo_alloc::allow(|| {
Arc::new(Body {
channel: channel.to_vec(),
payload: payload.to_vec(),
})
})
})
}
fn introspect(server: &Server, args: Args<'_>, out: &mut Out) -> Result<()> {
let sub = args.get(1);
if sub.eq_ignore_ascii_case(b"CHANNELS") {
names(server, args, out, Kind::Channel)
} else if sub.eq_ignore_ascii_case(b"SHARDCHANNELS") {
names(server, args, out, Kind::Shard)
} else if sub.eq_ignore_ascii_case(b"NUMSUB") {
counts(server, args, out, Kind::Channel);
Ok(())
} else if sub.eq_ignore_ascii_case(b"SHARDNUMSUB") {
counts(server, args, out, Kind::Shard);
Ok(())
} else if sub.eq_ignore_ascii_case(b"NUMPAT") {
if args.len() != 2 {
return Err(args::wrong_arity_sub("pubsub", "numpat"));
}
out.uint(server.pubsub.lock().patterns.len() as u64);
Ok(())
} else if sub.eq_ignore_ascii_case(b"HELP") {
if args.len() != 2 {
return Err(args::wrong_arity_sub("pubsub", "help"));
}
super::server::help(out, PUBSUB_HELP);
Ok(())
} else {
Err(args::unknown_subcommand(sub, "PUBSUB"))
}
}
fn names(server: &Server, args: Args<'_>, out: &mut Out, kind: Kind) -> Result<()> {
if args.len() > 3 {
return Err(unknown_or_arity(args.get(1), "PUBSUB"));
}
let pattern = args.opt(2);
let reg = server.pubsub.lock();
let table = reg.table(kind);
let hit = |name: &Vec<u8>| pattern.is_none_or(|p| glob::matches(p, name));
out.array(table.keys().filter(|name| hit(name)).count());
for name in table.keys().filter(|name| hit(name)) {
out.bulk(name);
}
Ok(())
}
fn counts(server: &Server, args: Args<'_>, out: &mut Out, kind: Kind) {
let reg = server.pubsub.lock();
out.array((args.len() - 2) * 2);
for i in 2..args.len() {
let name = args.get(i);
out.bulk(name);
out.uint(reg.table(kind).get(name).map_or(0, Vec::len) as u64);
}
}
fn unknown_or_arity(sub: &[u8], container: &str) -> Error {
yo_alloc::allow(|| {
let sub = String::from_utf8_lossy(sub);
Error::new(
yo_common::Code::Invalid,
format!(
"unknown subcommand or wrong number of arguments for '{sub}'. Try {container} HELP."
),
)
})
}
pub(crate) fn release(server: &Server, session: &mut Session) {
let Some(subs) = session.subs.take() else {
return;
};
if subs.total() == 0 {
return;
}
let client = session.id();
let thread = server.my_slot();
let mut reg = server.pubsub.lock();
for kind in KINDS {
for name in subs.list(kind) {
reg.remove(kind, name, client);
}
}
reg.clients -= 1;
server.note_here(thread, -1);
server.note_subs(®);
}
fn allowed(name: &str) -> bool {
matches!(
name,
"subscribe"
| "unsubscribe"
| "psubscribe"
| "punsubscribe"
| "ssubscribe"
| "sunsubscribe"
| "ping"
| "quit"
| "reset"
)
}
pub(crate) fn refused(session: &Session, spec: &Spec, out: &Out) -> Option<Error> {
if out.proto().is_resp3() || session.running() || session.scripted() {
return None;
}
if !session.subscribed() || allowed(spec.name) {
return None;
}
Some(yo_alloc::allow(|| {
let name = spec.name;
Error::new(
yo_common::Code::Invalid,
format!(
"Can't execute '{name}': only (P|S)SUBSCRIBE / (P|S)UNSUBSCRIBE / \
PING / QUIT / RESET are allowed in this context"
),
)
}))
}
pub(crate) fn ping(session: &Session, args: Args<'_>, out: &mut Out) -> bool {
if out.proto().is_resp3() || session.running() || session.scripted() || !session.subscribed() {
return false;
}
out.array(2);
out.bulk(b"pong");
out.bulk(args.opt(1).unwrap_or(b""));
true
}
const PUBSUB_HELP: &[&str] = &[
"PUBSUB <subcommand> [<arg> [value] [opt] ...]. Subcommands are:",
"CHANNELS [<pattern>]",
" Return the currently active channels matching a <pattern> (default: '*').",
"NUMPAT",
" Return number of subscriptions to patterns.",
"NUMSUB [<channel> ...]",
" Return the number of subscribers for the specified channels, excluding",
" pattern subscriptions(default: no channels).",
"SHARDCHANNELS [<pattern>]",
" Return the currently active shard level channels matching a <pattern> (default: '*').",
"SHARDNUMSUB [<shardchannel> ...]",
" Return the number of subscribers for the specified shard level channel(s)",
"HELP",
" Print this help.",
];