use crate::{
ClientError, Error, PubSubReceiver, Result,
client::{Client, ClientPreparedCommand},
commands::InternalPubSubCommands,
network::{PubSubPush, PubSubSender},
resp::{CommandArgs, CommandArgsMut, RefBulkString, RespResponse},
};
use bytes::Bytes;
use futures_util::{Stream, StreamExt};
use serde::Serialize;
use std::{
collections::HashSet,
fmt,
pin::Pin,
task::{Context, Poll},
};
use tracing::warn;
pub struct PubSubMessage {
buf: Box<[u8]>,
channel_start: usize,
payload_start: usize,
}
impl PubSubMessage {
#[inline]
pub fn pattern(&self) -> &[u8] {
&self.buf[..self.channel_start]
}
#[inline]
pub fn channel(&self) -> &[u8] {
&self.buf[self.channel_start..self.payload_start]
}
#[inline]
pub fn payload(&self) -> &[u8] {
&self.buf[self.payload_start..]
}
#[inline]
fn from_segments(pattern: &[u8], channel: &[u8], payload: &[u8]) -> Self {
let channel_start = pattern.len();
let payload_start = channel_start.saturating_add(channel.len());
let mut buf = Vec::with_capacity(payload_start.saturating_add(payload.len()));
buf.extend_from_slice(pattern);
buf.extend_from_slice(channel);
buf.extend_from_slice(payload);
Self {
buf: buf.into_boxed_slice(),
channel_start,
payload_start,
}
}
}
impl TryFrom<&RespResponse> for PubSubMessage {
type Error = Error;
#[inline]
fn try_from(response: &RespResponse) -> Result<Self> {
match PubSubPush::try_from(response) {
Ok(PubSubPush::Message(channel, payload) | PubSubPush::SMessage(channel, payload)) => {
Ok(Self::from_segments(&[], channel, payload))
}
Ok(PubSubPush::PMessage(pattern, channel, payload)) => {
Ok(Self::from_segments(pattern, channel, payload))
}
_ => Err(Error::from(ClientError::UnexpectedPubSubMessage)),
}
}
}
impl fmt::Debug for PubSubMessage {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PubSubMessage")
.field("pattern", &String::from_utf8_lossy(self.pattern()))
.field("channel", &String::from_utf8_lossy(self.channel()))
.field("payload", &String::from_utf8_lossy(self.payload()))
.finish()
}
}
fn extract_args_to_set(args: CommandArgs, set: &mut HashSet<Bytes>) {
for arg in &args {
set.insert(arg);
}
}
pub struct PubSubSplitSink {
closed: bool,
channels: HashSet<Bytes>,
patterns: HashSet<Bytes>,
shardchannels: HashSet<Bytes>,
sender: PubSubSender,
client: Client,
}
impl PubSubSplitSink {
pub async fn subscribe(&mut self, channels: impl Serialize) -> Result<()> {
let channels = CommandArgsMut::default().arg(channels).freeze();
for channel in &channels {
if self.channels.contains(&channel) {
return Err(Error::from(ClientError::AlreadySubscribed));
}
}
self.client
.subscribe_from_pub_sub_sender(&channels, &self.sender)
.await?;
extract_args_to_set(channels, &mut self.channels);
Ok(())
}
pub async fn psubscribe(&mut self, patterns: impl Serialize) -> Result<()> {
let patterns = CommandArgsMut::default().arg(patterns).freeze();
for pattern in &patterns {
if self.patterns.contains(&pattern) {
return Err(Error::from(ClientError::AlreadySubscribed));
}
}
self.client
.psubscribe_from_pub_sub_sender(&patterns, &self.sender)
.await?;
extract_args_to_set(patterns, &mut self.patterns);
Ok(())
}
pub async fn ssubscribe(&mut self, shardchannels: impl Serialize) -> Result<()> {
let shardchannels = CommandArgsMut::default().arg(shardchannels).freeze();
for shardchannel in &shardchannels {
if self.shardchannels.contains(&shardchannel) {
return Err(Error::from(ClientError::AlreadySubscribed));
}
}
self.client
.ssubscribe_from_pub_sub_sender(&shardchannels, &self.sender)
.await?;
extract_args_to_set(shardchannels, &mut self.shardchannels);
Ok(())
}
pub async fn unsubscribe(&mut self, channels: impl Serialize) -> Result<()> {
let channels = CommandArgsMut::default().arg(channels).freeze();
self.client.unsubscribe(&channels).await?;
for channel in &channels {
self.channels.remove(&channel);
}
Ok(())
}
pub async fn punsubscribe(&mut self, patterns: impl Serialize) -> Result<()> {
let patterns = CommandArgsMut::default().arg(patterns).freeze();
self.client.punsubscribe(&patterns).await?;
for pattern in &patterns {
self.patterns.remove(&pattern);
}
Ok(())
}
pub async fn sunsubscribe(&mut self, shardchannels: impl Serialize) -> Result<()> {
let shardchannels = CommandArgsMut::default().arg(shardchannels).freeze();
self.client.sunsubscribe(&shardchannels).await?;
for shardchannel in &shardchannels {
self.shardchannels.remove(&shardchannel);
}
Ok(())
}
pub async fn close(mut self) -> Result<()> {
if self.closed {
return Ok(());
}
if !self.channels.is_empty() {
let mut args = CommandArgsMut::default();
for channel in &self.channels {
args = args.arg(channel);
}
self.client.unsubscribe(args).await?;
self.channels.clear();
}
if !self.patterns.is_empty() {
let mut args = CommandArgsMut::default();
for pattern in &self.patterns {
args = args.arg(pattern);
}
self.client.punsubscribe(args).await?;
self.patterns.clear();
}
if !self.shardchannels.is_empty() {
let mut args = CommandArgsMut::default();
for shardchannel in &self.shardchannels {
args = args.arg(shardchannel);
}
self.client.sunsubscribe(args).await?;
self.shardchannels.clear();
}
self.closed = true;
Ok(())
}
}
impl Drop for PubSubSplitSink {
fn drop(&mut self) {
if self.closed {
return;
}
if !self.channels.is_empty() {
let mut args = CommandArgsMut::default();
for channel in &self.channels {
args = args.arg(RefBulkString::new(channel.as_ref()));
}
if let Err(e) = self.client.unsubscribe(args).forget() {
warn!("Error while unsubscribing from the dropped channels: {e}");
}
self.channels.clear();
}
if !self.patterns.is_empty() {
let mut args = CommandArgsMut::default();
for pattern in &self.patterns {
args = args.arg(RefBulkString::new(pattern.as_ref()));
}
if let Err(e) = self.client.punsubscribe(args).forget() {
warn!("Error while unsubscribing from the dropped patterns: {e}");
}
self.patterns.clear();
}
if !self.shardchannels.is_empty() {
let mut args = CommandArgsMut::default();
for shardchannel in &self.shardchannels {
args = args.arg(RefBulkString::new(shardchannel.as_ref()));
}
if let Err(e) = self.client.sunsubscribe(args).forget() {
warn!("Error while unsubscribing from the dropped shard channels: {e}");
}
self.shardchannels.clear();
}
self.closed = true;
}
}
pub struct PubSubSplitStream {
receiver: PubSubReceiver,
}
impl PubSubSplitStream {
pub fn dropped_messages(&self) -> usize {
self.receiver.dropped_messages()
}
}
impl Stream for PubSubSplitStream {
type Item = Result<PubSubMessage>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
match self.get_mut().receiver.poll_next_unpin(cx) {
Poll::Ready(Some(Ok(response))) => {
Poll::Ready(Some(PubSubMessage::try_from(&response)))
}
Poll::Ready(None) => Poll::Ready(None),
Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(e))),
Poll::Pending => Poll::Pending,
}
}
}
pub struct PubSubStream {
split_sink: PubSubSplitSink,
split_stream: PubSubSplitStream,
}
impl PubSubStream {
pub fn dropped_messages(&self) -> usize {
self.split_stream.dropped_messages()
}
pub(crate) fn new(sender: PubSubSender, receiver: PubSubReceiver, client: Client) -> Self {
Self {
split_sink: PubSubSplitSink {
closed: false,
channels: HashSet::default(),
patterns: HashSet::default(),
shardchannels: HashSet::default(),
sender,
client,
},
split_stream: PubSubSplitStream { receiver },
}
}
pub(crate) fn from_channels(
channels: CommandArgs,
sender: PubSubSender,
receiver: PubSubReceiver,
client: Client,
) -> Self {
let mut set = HashSet::with_capacity(channels.len());
extract_args_to_set(channels, &mut set);
Self {
split_sink: PubSubSplitSink {
closed: false,
channels: set,
patterns: HashSet::default(),
shardchannels: HashSet::default(),
sender,
client,
},
split_stream: PubSubSplitStream { receiver },
}
}
pub(crate) fn from_patterns(
patterns: CommandArgs,
sender: PubSubSender,
receiver: PubSubReceiver,
client: Client,
) -> Self {
let mut set: HashSet<Bytes> = HashSet::with_capacity(patterns.len());
extract_args_to_set(patterns, &mut set);
Self {
split_sink: PubSubSplitSink {
closed: false,
channels: HashSet::default(),
patterns: set,
shardchannels: HashSet::default(),
sender,
client,
},
split_stream: PubSubSplitStream { receiver },
}
}
pub(crate) fn from_shardchannels(
shardchannels: CommandArgs,
sender: PubSubSender,
receiver: PubSubReceiver,
client: Client,
) -> Self {
let mut set: HashSet<Bytes> = HashSet::with_capacity(shardchannels.len());
extract_args_to_set(shardchannels, &mut set);
Self {
split_sink: PubSubSplitSink {
closed: false,
channels: HashSet::default(),
patterns: HashSet::default(),
shardchannels: set,
sender,
client,
},
split_stream: PubSubSplitStream { receiver },
}
}
pub async fn subscribe(&mut self, channels: impl Serialize) -> Result<()> {
self.split_sink.subscribe(channels).await
}
pub async fn psubscribe(&mut self, patterns: impl Serialize) -> Result<()> {
self.split_sink.psubscribe(patterns).await
}
pub async fn ssubscribe(&mut self, shardchannels: impl Serialize) -> Result<()> {
self.split_sink.ssubscribe(shardchannels).await
}
pub async fn unsubscribe(&mut self, channels: impl Serialize) -> Result<()> {
self.split_sink.unsubscribe(channels).await
}
pub async fn punsubscribe(&mut self, patterns: impl Serialize) -> Result<()> {
self.split_sink.punsubscribe(patterns).await
}
pub async fn sunsubscribe(&mut self, shardchannels: impl Serialize) -> Result<()> {
self.split_sink.sunsubscribe(shardchannels).await
}
pub fn split(self) -> (PubSubSplitSink, PubSubSplitStream) {
(self.split_sink, self.split_stream)
}
pub async fn close(self) -> Result<()> {
self.split_sink.close().await
}
}
impl Stream for PubSubStream {
type Item = Result<PubSubMessage>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
if self.split_sink.closed {
Poll::Ready(None)
} else {
let pinned = std::pin::pin!(&mut self.get_mut().split_stream);
pinned.poll_next(cx)
}
}
}