use redis::AsyncCommands;
use serde::Serialize;
use serde::de::DeserializeOwned;
use tokio::sync::broadcast;
use tracing::{debug, trace, warn};
use crate::error::MajraError;
pub struct RedisPubSub {
client: redis::Client,
prefix: String,
}
impl RedisPubSub {
pub fn new(client: redis::Client, prefix: impl Into<String>) -> Self {
Self {
client,
prefix: prefix.into(),
}
}
pub async fn publish<T: Serialize>(
&self,
topic: &str,
payload: &T,
) -> Result<usize, MajraError> {
let channel = format!("{}{}", self.prefix, topic);
let data = serde_json::to_string(payload).map_err(|e| MajraError::PubSub(e.to_string()))?;
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::PubSub(format!("redis connection: {e}")))?;
let receivers: usize = conn
.publish(&channel, &data)
.await
.map_err(|e| MajraError::PubSub(format!("redis publish: {e}")))?;
trace!(channel = channel.as_str(), receivers, "redis: published");
Ok(receivers)
}
pub fn subscribe<T: DeserializeOwned + Clone + Send + 'static>(
&self,
pattern: &str,
capacity: usize,
) -> Result<broadcast::Receiver<(String, T)>, MajraError> {
let (tx, rx) = broadcast::channel(capacity);
let channel_pattern = format!("{}{}", self.prefix, pattern);
let client = self.client.clone();
let prefix_len = self.prefix.len();
tokio::spawn(async move {
let mut pubsub_conn = match client.get_async_pubsub().await {
Ok(c) => c,
Err(e) => {
warn!(error = %e, "redis subscribe: connection failed");
return;
}
};
if let Err(e) = pubsub_conn.psubscribe(&channel_pattern).await {
warn!(error = %e, "redis subscribe: psubscribe failed");
return;
}
debug!(pattern = channel_pattern.as_str(), "redis: subscribed");
use futures_util::StreamExt;
let mut msg_stream = pubsub_conn.on_message();
while let Some(msg) = msg_stream.next().await {
let channel: String = match msg.get_channel() {
Ok(c) => c,
Err(_) => continue,
};
let payload_str: String = match msg.get_payload() {
Ok(p) => p,
Err(_) => continue,
};
let topic = if channel.len() > prefix_len {
channel[prefix_len..].to_string()
} else {
channel.clone()
};
match serde_json::from_str::<T>(&payload_str) {
Ok(value) => {
if tx.send((topic, value)).is_err() {
break;
}
}
Err(e) => {
trace!(error = %e, "redis: failed to deserialize message");
}
}
}
});
Ok(rx)
}
#[inline]
pub fn prefix(&self) -> &str {
&self.prefix
}
}
pub struct RedisQueue {
client: redis::Client,
key: String,
}
impl RedisQueue {
pub fn new(client: redis::Client, key: impl Into<String>) -> Self {
Self {
client,
key: key.into(),
}
}
pub async fn enqueue<T: Serialize>(
&self,
priority: u8,
payload: &T,
) -> Result<usize, MajraError> {
let data = serde_json::to_string(payload).map_err(|e| MajraError::Queue(e.to_string()))?;
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Queue(format!("redis connection: {e}")))?;
let score = -(i64::from(priority));
conn.zadd::<_, _, _, ()>(&self.key, &data, score)
.await
.map_err(|e| MajraError::Queue(format!("redis zadd: {e}")))?;
let len: usize = conn
.zcard(&self.key)
.await
.map_err(|e| MajraError::Queue(format!("redis zcard: {e}")))?;
trace!(priority, len, "redis queue: enqueued");
Ok(len)
}
pub async fn dequeue<T: DeserializeOwned>(&self) -> Result<Option<T>, MajraError> {
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Queue(format!("redis connection: {e}")))?;
let result: Vec<(String, f64)> = redis::cmd("ZPOPMIN")
.arg(&self.key)
.arg(1)
.query_async(&mut conn)
.await
.map_err(|e| MajraError::Queue(format!("redis zpopmin: {e}")))?;
match result.first() {
Some((data, _score)) => {
let value: T = serde_json::from_str(data)
.map_err(|e| MajraError::Queue(format!("deserialize: {e}")))?;
trace!("redis queue: dequeued");
Ok(Some(value))
}
None => Ok(None),
}
}
pub async fn len(&self) -> Result<usize, MajraError> {
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Queue(format!("redis connection: {e}")))?;
let len: usize = conn
.zcard(&self.key)
.await
.map_err(|e| MajraError::Queue(format!("redis zcard: {e}")))?;
Ok(len)
}
pub async fn is_empty(&self) -> Result<bool, MajraError> {
self.len().await.map(|n| n == 0)
}
pub async fn peek<T: DeserializeOwned>(&self) -> Result<Option<T>, MajraError> {
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Queue(format!("redis connection: {e}")))?;
let result: Vec<(String, f64)> = conn
.zrangebyscore_limit_withscores(&self.key, "-inf", "+inf", 0, 1)
.await
.map_err(|e| MajraError::Queue(format!("redis zrangebyscore: {e}")))?;
match result.first() {
Some((data, _)) => {
let value: T = serde_json::from_str(data)
.map_err(|e| MajraError::Queue(format!("deserialize: {e}")))?;
Ok(Some(value))
}
None => Ok(None),
}
}
pub async fn clear(&self) -> Result<(), MajraError> {
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Queue(format!("redis connection: {e}")))?;
conn.del::<_, ()>(&self.key)
.await
.map_err(|e| MajraError::Queue(format!("redis del: {e}")))?;
debug!(key = self.key.as_str(), "redis queue: cleared");
Ok(())
}
#[inline]
pub fn key(&self) -> &str {
&self.key
}
}
pub struct RedisRateLimiter {
client: redis::Client,
rate: f64,
burst: usize,
prefix: String,
script: redis::Script,
}
const RATE_LIMIT_LUA: &str = r#"
local key = KEYS[1]
local rate = tonumber(ARGV[1])
local burst = tonumber(ARGV[2])
local now_ms = tonumber(ARGV[3])
local tokens = tonumber(redis.call('HGET', key, 'tokens') or burst)
local last_ms = tonumber(redis.call('HGET', key, 'last_ms') or now_ms)
local elapsed_s = (now_ms - last_ms) / 1000.0
tokens = math.min(tokens + elapsed_s * rate, burst)
if tokens >= 1 then
tokens = tokens - 1
redis.call('HSET', key, 'tokens', tostring(tokens), 'last_ms', tostring(now_ms))
redis.call('EXPIRE', key, math.ceil(burst / rate) + 60)
return 1
else
redis.call('HSET', key, 'tokens', tostring(tokens), 'last_ms', tostring(now_ms))
redis.call('EXPIRE', key, math.ceil(burst / rate) + 60)
return 0
end
"#;
impl RedisRateLimiter {
pub fn new(client: redis::Client, rate: f64, burst: usize, prefix: impl Into<String>) -> Self {
Self {
client,
rate,
burst,
prefix: prefix.into(),
script: redis::Script::new(RATE_LIMIT_LUA),
}
}
pub async fn check(&self, key: &str) -> crate::error::Result<bool> {
let redis_key = format!("{}{key}", self.prefix);
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Queue(e.to_string()))?;
let now_ms = chrono::Utc::now().timestamp_millis();
let result: i32 = self
.script
.key(&redis_key)
.arg(self.rate)
.arg(self.burst)
.arg(now_ms)
.invoke_async(&mut conn)
.await
.map_err(|e| MajraError::Queue(e.to_string()))?;
trace!(key, allowed = result == 1, "redis-rl: check");
Ok(result == 1)
}
#[inline]
pub fn rate(&self) -> f64 {
self.rate
}
#[inline]
pub fn burst(&self) -> usize {
self.burst
}
#[inline]
pub fn prefix(&self) -> &str {
&self.prefix
}
}
pub struct RedisHeartbeatTracker {
client: redis::Client,
prefix: String,
ttl_secs: u64,
}
impl RedisHeartbeatTracker {
pub fn new(client: redis::Client, prefix: impl Into<String>, ttl_secs: u64) -> Self {
Self {
client,
prefix: prefix.into(),
ttl_secs,
}
}
pub async fn register(
&self,
node_id: &str,
metadata: &serde_json::Value,
) -> crate::error::Result<()> {
let key = format!("{}{node_id}", self.prefix);
let value =
serde_json::to_string(metadata).map_err(|e| MajraError::Heartbeat(e.to_string()))?;
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
conn.set_ex::<_, _, ()>(&key, &value, self.ttl_secs)
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
debug!(node_id, ttl = self.ttl_secs, "redis-hb: registered");
Ok(())
}
pub async fn heartbeat(
&self,
node_id: &str,
metadata: &serde_json::Value,
) -> crate::error::Result<()> {
self.register(node_id, metadata).await
}
pub async fn is_online(&self, node_id: &str) -> crate::error::Result<bool> {
let key = format!("{}{node_id}", self.prefix);
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
let exists: bool = conn
.exists(&key)
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
Ok(exists)
}
pub async fn get_metadata(
&self,
node_id: &str,
) -> crate::error::Result<Option<serde_json::Value>> {
let key = format!("{}{node_id}", self.prefix);
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
let value: Option<String> = conn
.get(&key)
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
match value {
Some(v) => {
let parsed: serde_json::Value =
serde_json::from_str(&v).map_err(|e| MajraError::Heartbeat(e.to_string()))?;
Ok(Some(parsed))
}
None => Ok(None),
}
}
pub async fn list_online(&self) -> crate::error::Result<Vec<(String, serde_json::Value)>> {
let pattern = format!("{}*", self.prefix);
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
let keys: Vec<String> = redis::cmd("KEYS")
.arg(&pattern)
.query_async(&mut conn)
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
let mut result = Vec::with_capacity(keys.len());
for key in &keys {
let value: Option<String> = conn
.get(key)
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
if let Some(v) = value {
let node_id = key.strip_prefix(&self.prefix).unwrap_or(key);
let metadata: serde_json::Value =
serde_json::from_str(&v).map_err(|e| MajraError::Heartbeat(e.to_string()))?;
result.push((node_id.to_string(), metadata));
}
}
Ok(result)
}
pub async fn deregister(&self, node_id: &str) -> crate::error::Result<()> {
let key = format!("{}{node_id}", self.prefix);
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
conn.del::<_, ()>(&key)
.await
.map_err(|e| MajraError::Heartbeat(e.to_string()))?;
debug!(node_id, "redis-hb: deregistered");
Ok(())
}
#[inline]
pub fn ttl_secs(&self) -> u64 {
self.ttl_secs
}
#[inline]
pub fn prefix(&self) -> &str {
&self.prefix
}
}
#[cfg(test)]
mod tests {
use super::*;
fn test_client() -> redis::Client {
redis::Client::open("redis://127.0.0.1/").expect("redis client")
}
#[tokio::test]
#[ignore]
async fn queue_enqueue_dequeue() {
let client = test_client();
let q = RedisQueue::new(client, format!("majra:test:queue:{}", uuid::Uuid::new_v4()));
q.enqueue(2u8, &serde_json::json!({"job": "normal"}))
.await
.unwrap();
q.enqueue(4u8, &serde_json::json!({"job": "critical"}))
.await
.unwrap();
q.enqueue(0u8, &serde_json::json!({"job": "background"}))
.await
.unwrap();
assert_eq!(q.len().await.unwrap(), 3);
let first: serde_json::Value = q.dequeue().await.unwrap().unwrap();
assert_eq!(first["job"], "critical");
let second: serde_json::Value = q.dequeue().await.unwrap().unwrap();
assert_eq!(second["job"], "normal");
let third: serde_json::Value = q.dequeue().await.unwrap().unwrap();
assert_eq!(third["job"], "background");
assert!(q.is_empty().await.unwrap());
q.clear().await.unwrap();
}
#[tokio::test]
#[ignore]
async fn queue_dequeue_empty() {
let client = test_client();
let q = RedisQueue::new(client, format!("majra:test:queue:{}", uuid::Uuid::new_v4()));
let result: Option<serde_json::Value> = q.dequeue().await.unwrap();
assert!(result.is_none());
}
#[tokio::test]
#[ignore]
async fn queue_peek() {
let client = test_client();
let q = RedisQueue::new(client, format!("majra:test:queue:{}", uuid::Uuid::new_v4()));
q.enqueue(3u8, &serde_json::json!({"job": "high"}))
.await
.unwrap();
let peeked: serde_json::Value = q.peek().await.unwrap().unwrap();
assert_eq!(peeked["job"], "high");
assert_eq!(q.len().await.unwrap(), 1);
q.clear().await.unwrap();
}
#[tokio::test]
#[ignore]
async fn pubsub_publish() {
let client = test_client();
let hub = RedisPubSub::new(client, "majra:test:ps:");
let receivers = hub
.publish("events/test", &serde_json::json!({"hello": "world"}))
.await
.unwrap();
assert_eq!(receivers, 0);
}
#[test]
fn pubsub_prefix() {
let client = test_client();
let hub = RedisPubSub::new(client, "majra:");
assert_eq!(hub.prefix(), "majra:");
}
#[test]
fn queue_key() {
let client = test_client();
let q = RedisQueue::new(client, "majra:queue:jobs");
assert_eq!(q.key(), "majra:queue:jobs");
}
#[test]
fn ratelimiter_config() {
let client = test_client();
let rl = RedisRateLimiter::new(client, 10.0, 100, "majra:rl:");
assert_eq!(rl.rate(), 10.0);
assert_eq!(rl.burst(), 100);
assert_eq!(rl.prefix(), "majra:rl:");
}
#[test]
fn heartbeat_tracker_config() {
let client = test_client();
let hb = RedisHeartbeatTracker::new(client, "majra:hb:", 30);
assert_eq!(hb.ttl_secs(), 30);
assert_eq!(hb.prefix(), "majra:hb:");
}
}