use redis::ControlFlow;
use redis::{Client, Connection, IntoConnectionInfo, Msg, ToRedisArgs};
use std::fmt::Debug;
use super::retry::Retry;
#[derive(Debug)]
pub enum RedisPubSubError {
AddressAlreadyConsumed,
FailToOpenClient,
FailToGetConnection,
FailToSubscribeToChannels,
FailToSubscribeToChannelsWithPatterns,
FailToGetMessage,
}
pub struct RedisPubSub<I, A>
where
I: IntoConnectionInfo + Sized + Debug + Clone,
A: ToRedisArgs,
{
addr: Option<I>,
channels: Vec<A>,
patterns: Vec<A>,
retry: Retry,
}
impl<I, A> RedisPubSub<I, A>
where
I: IntoConnectionInfo + Sized + Debug + Clone,
A: ToRedisArgs,
{
pub fn new(addr: I) -> Self {
Self {
addr: Some(addr),
channels: Vec::new(),
patterns: Vec::new(),
retry: Retry::new(),
}
}
pub fn set_retry(&mut self, max_count: u32, init_delay_ms: u64, max_delay_ms: u64) {
self.retry = Retry::with_params(max_count, init_delay_ms, max_delay_ms);
}
pub fn subscribe(&mut self, channels: A) {
self.channels.push(channels);
}
pub fn psubscribe(&mut self, patterns: A) {
self.patterns.push(patterns);
}
pub fn receive<F, U>(mut self, mut f: F) -> errs::Result<U>
where
F: FnMut(Msg) -> ControlFlow<U>,
{
let addr = self
.addr
.take()
.ok_or_else(|| errs::Err::new(RedisPubSubError::AddressAlreadyConsumed))?;
let client = Client::open(addr)
.map_err(|e| errs::Err::with_source(RedisPubSubError::FailToOpenClient, e))?;
loop {
let mut conn: Connection = match client.get_connection() {
Ok(c) => c,
Err(e) => {
if self.retry.wait_with_backoff() {
continue;
}
return Err(errs::Err::with_source(
RedisPubSubError::FailToGetConnection,
e,
));
}
};
let mut pubsub = conn.as_pubsub();
for c in self.channels.iter() {
pubsub.subscribe(c).map_err(|e| {
errs::Err::with_source(RedisPubSubError::FailToSubscribeToChannels, e)
})?;
}
for p in self.patterns.iter() {
pubsub.psubscribe(p).map_err(|e| {
errs::Err::with_source(
RedisPubSubError::FailToSubscribeToChannelsWithPatterns,
e,
)
})?;
}
loop {
match pubsub.get_message() {
Ok(msg) => {
self.retry.reset();
if let ControlFlow::Break(value) = f(msg) {
return Ok(value);
}
}
Err(e) => {
if self.retry.wait_with_backoff() {
continue;
}
return Err(errs::Err::with_source(
RedisPubSubError::FailToGetMessage,
e,
));
}
}
}
}
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
use crate::pubsub::{RedisPubSubMsgDataConn, RedisPubSubMsgDataSrc};
use crate::standalone_sync::{RedisDataConn, RedisDataSrc};
use override_macro::{overridable, override_with};
use redis::{ControlFlow, TypedCommands};
use sabi::{DataAcc, DataHub};
#[overridable]
trait PublishData {
fn say_greet(&mut self, s: &str) -> errs::Result<()>;
}
fn publish_logic(data: &mut impl PublishData) -> errs::Result<()> {
data.say_greet("Hello")?;
Ok(())
}
#[overridable]
trait SubscribeData {
fn receive_greet(&mut self) -> errs::Result<String>;
}
fn subscribe_logic(data: &mut impl SubscribeData) -> errs::Result<()> {
assert_eq!(data.receive_greet()?, "Hello");
Ok(())
}
#[overridable]
trait RedisPubSubDataAcc: DataAcc {
fn say_greet(&mut self, s: &str) -> errs::Result<()> {
let data_conn = self.get_data_conn::<RedisDataConn>("redis")?;
let conn = data_conn.get_connection();
std::thread::sleep(std::time::Duration::from_millis(100));
conn.publish("channel-1", s).unwrap();
Ok(())
}
fn receive_greet(&mut self) -> errs::Result<String> {
let data_conn = self.get_data_conn::<RedisPubSubMsgDataConn>("redis/pubsub")?;
let msg = data_conn.get_message();
let payload: String = msg.get_payload().unwrap();
Ok(payload)
}
}
impl RedisPubSubDataAcc for DataHub {}
#[override_with(RedisPubSubDataAcc)]
impl PublishData for DataHub {}
#[override_with(RedisPubSubDataAcc)]
impl SubscribeData for DataHub {}
#[test]
fn test() -> errs::Result<()> {
{
let _ = std::thread::spawn(move || {
let mut data = DataHub::new();
data.uses("redis", RedisDataSrc::new("redis://127.0.0.1/"));
data.run(publish_logic).unwrap();
});
}
{
let mut pubsub = RedisPubSub::new("redis://127.0.0.1/");
pubsub.subscribe("channel-1");
let n = pubsub.receive(|msg| {
let mut data = DataHub::new();
data.uses("redis/pubsub", RedisPubSubMsgDataSrc::new(msg));
data.run(subscribe_logic).unwrap();
ControlFlow::Break(1)
})?;
assert_eq!(n, 1);
}
Ok(())
}
}