use super::retry::Retry;
use futures::{stream::StreamExt, Future};
use redis::aio::PubSub;
use redis::{Client, ControlFlow, IntoConnectionInfo, Msg, ToRedisArgs};
use std::fmt::Debug;
#[derive(Debug)]
pub enum RedisPubSubAsyncError {
AddressAlreadyConsumed,
FailToOpenClient,
FailToGetAsyncPubSub,
FailToSubscribeToChannels,
FailToSubscribeToChannelsWithPatterns,
FailToGetMessage,
}
pub struct RedisPubSubAsync<I, A>
where
I: IntoConnectionInfo + Sized + Debug + Clone,
A: ToRedisArgs,
{
addr: Option<I>,
channels: Vec<A>,
patterns: Vec<A>,
retry: Retry,
}
impl<I, A> RedisPubSubAsync<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 async fn receive_async<F, Fut, U>(mut self, mut f: F) -> errs::Result<U>
where
F: FnMut(Msg) -> Fut,
Fut: Future<Output = ControlFlow<U>>,
{
let addr = self
.addr
.take()
.ok_or_else(|| errs::Err::new(RedisPubSubAsyncError::AddressAlreadyConsumed))?;
let client = Client::open(addr)
.map_err(|e| errs::Err::with_source(RedisPubSubAsyncError::FailToOpenClient, e))?;
loop {
let pubsub: PubSub = match client.get_async_pubsub().await {
Ok(pubsub) => pubsub,
Err(e) => {
if self.retry.wait_with_backoff_async().await {
continue;
}
return Err(errs::Err::with_source(
RedisPubSubAsyncError::FailToGetAsyncPubSub,
e,
));
}
};
let (mut sink, mut stream) = pubsub.split();
for c in self.channels.iter() {
sink.subscribe(c).await.map_err(|e| {
errs::Err::with_source(RedisPubSubAsyncError::FailToSubscribeToChannels, e)
})?;
}
for p in self.patterns.iter() {
sink.psubscribe(p).await.map_err(|e| {
errs::Err::with_source(
RedisPubSubAsyncError::FailToSubscribeToChannelsWithPatterns,
e,
)
})?;
}
loop {
match stream.next().await {
Some(msg) => {
self.retry.reset();
if let ControlFlow::Break(value) = f(msg).await {
return Ok(value);
}
}
None => {
if self.retry.wait_with_backoff_async().await {
continue;
}
return Err(errs::Err::new(RedisPubSubAsyncError::FailToGetMessage));
}
}
}
}
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
use crate::pubsub::{RedisPubSubMsgAsyncDataConn, RedisPubSubMsgAsyncDataSrc};
use crate::standalone_async::{RedisAsyncDataConn, RedisAsyncDataSrc};
use override_macro::{overridable, override_with};
use redis::{AsyncTypedCommands, ControlFlow};
use sabi::tokio::{logic, DataAcc, DataHub};
use tokio::time;
#[overridable]
trait PublishData {
async fn say_greet_async(&mut self, s: &str) -> errs::Result<()>;
}
async fn publish_logic_async(data: &mut impl PublishData) -> errs::Result<()> {
data.say_greet_async("Hello").await?;
Ok(())
}
#[overridable]
trait SubscribeData {
async fn receive_greet_async(&mut self) -> errs::Result<String>;
}
async fn subscribe_logic_async(data: &mut impl SubscribeData) -> errs::Result<()> {
assert_eq!(data.receive_greet_async().await?, "Hello");
Ok(())
}
#[overridable]
trait RedisPubSubAsyncDataAcc: DataAcc {
async fn say_greet_async(&mut self, s: &str) -> errs::Result<()> {
let data_conn = self
.get_data_conn_async::<RedisAsyncDataConn>("redis")
.await?;
let conn = data_conn.get_connection();
time::sleep(time::Duration::from_millis(100)).await;
conn.publish("channel-1", s).await.unwrap();
Ok(())
}
async fn receive_greet_async(&mut self) -> errs::Result<String> {
let data_conn = self
.get_data_conn_async::<RedisPubSubMsgAsyncDataConn>("redis/pubsub")
.await?;
let msg = data_conn.get_message();
let payload: String = msg.get_payload().unwrap();
Ok(payload)
}
}
impl RedisPubSubAsyncDataAcc for DataHub {}
#[override_with(RedisPubSubAsyncDataAcc)]
impl PublishData for DataHub {}
#[override_with(RedisPubSubAsyncDataAcc)]
impl SubscribeData for DataHub {}
#[tokio::test]
async fn test() {
{
let _ = tokio::spawn(async {
let mut data = DataHub::new();
data.uses("redis", RedisAsyncDataSrc::new("redis://127.0.0.1:6379/8"));
data.run_async(logic!(publish_logic_async)).await.unwrap();
});
}
{
let mut pubsub = RedisPubSubAsync::new("redis://127.0.0.1:6379/8");
pubsub.subscribe("channel-1");
let n = pubsub
.receive_async(async |msg| {
let mut data = DataHub::new();
data.uses("redis/pubsub", RedisPubSubMsgAsyncDataSrc::new(msg));
data.run_async(logic!(subscribe_logic_async)).await.unwrap();
ControlFlow::Break(1)
})
.await
.unwrap();
assert_eq!(n, 1);
}
}
}