use super::{
JobIndexPage, JobIndexRow, JobIndexWaitStats, RangeItem, RangeOrder, RangePage, Storage,
StorageOperation,
};
use crate::UtcDateTime;
use redis::{aio::MultiplexedConnection, AsyncCommands};
const METRIC_BUCKET_MS: i64 = 60_000;
fn minute_bucket(at: UtcDateTime) -> i64 {
at.timestamp_millis().div_euclid(METRIC_BUCKET_MS)
}
const APPLY_SCRIPT: &str = r#"
local i = 1
while i <= #ARGV do
local operation = ARGV[i]
local key = ARGV[i + 1]
local value = ARGV[i + 2]
if operation == 'set' then
redis.call('SET', key, value)
elseif operation == 'delete' then
redis.call('DEL', key)
elseif operation == 'expire' then
redis.call('EXPIRE', key, tonumber(value))
elseif operation == 'range_add' then
if not redis.call('ZSCORE', key, value) then
local cursor = redis.call('INCR', key .. ':later-cursor')
redis.call('ZADD', key, cursor, value)
end
redis.call('ZREM', key .. ':later-expiry', value)
elseif operation == 'range_remove' then
redis.call('ZREM', key, value)
redis.call('ZREM', key .. ':later-expiry', value)
elseif operation == 'range_expire' then
local now = redis.call('TIME')
redis.call('ZADD', key .. ':later-expiry', now[1] + tonumber(ARGV[i + 3]), value)
i = i + 4
value = nil
elseif operation == 'range_clear' then
redis.call('DEL', key, key .. ':later-cursor', key .. ':later-expiry')
else
return redis.error_reply('unknown storage operation: ' .. operation)
end
if value then
i = i + 3
end
end
return 1
"#;
const REMOVE_EXPIRED_SCRIPT: &str = r#"
local now = redis.call('TIME')
local expiry_key = KEYS[1] .. ':later-expiry'
local expired = redis.call('ZRANGEBYSCORE', expiry_key, '-inf', now[1])
if #expired > 0 then
redis.call('ZREM', KEYS[1], unpack(expired))
redis.call('ZREM', expiry_key, unpack(expired))
end
return #expired
"#;
#[derive(Clone)]
pub struct Redis {
connection: MultiplexedConnection,
}
impl Redis {
pub async fn new(url: &str) -> anyhow::Result<Self> {
let client = redis::Client::open(url)?;
let connection = client.get_multiplexed_tokio_connection().await?;
Ok(Self { connection })
}
fn score_to_cursor(score: f64) -> anyhow::Result<i64> {
if !score.is_finite() || score.fract() != 0.0 {
return Err(anyhow::anyhow!("Redis returned an invalid range cursor"));
}
Ok(score.to_string().parse()?)
}
async fn remove_expired(&self, key: &str) -> anyhow::Result<()> {
let mut connection = self.connection.clone();
redis::Script::new(REMOVE_EXPIRED_SCRIPT)
.key(key)
.invoke_async::<_, usize>(&mut connection)
.await?;
Ok(())
}
}
#[async_trait::async_trait]
impl Storage for Redis {
async fn get(&self, key: &str) -> anyhow::Result<Option<Vec<u8>>> {
let mut connection = self.connection.clone();
Ok(connection.get(key).await?)
}
async fn set_if_absent(&self, key: &str, value: &[u8]) -> anyhow::Result<bool> {
let mut connection = self.connection.clone();
Ok(connection.set_nx(key, value).await?)
}
async fn apply(&self, operations: Vec<StorageOperation>) -> anyhow::Result<()> {
for operation in &operations {
match operation {
StorageOperation::Expire { ttl_seconds, .. } => {
i64::try_from(*ttl_seconds)?;
}
StorageOperation::Set { .. }
| StorageOperation::Delete { .. }
| StorageOperation::RangeAdd { .. }
| StorageOperation::RangeRemove { .. }
| StorageOperation::RangeClear { .. } => {}
StorageOperation::RangeExpire { ttl_seconds, .. } => {
i64::try_from(*ttl_seconds)?;
}
}
}
let script = redis::Script::new(APPLY_SCRIPT);
let mut invocation = script.prepare_invoke();
for operation in operations {
match operation {
StorageOperation::Set { key, value } => {
invocation.arg("set").arg(key).arg(value);
}
StorageOperation::Delete { key } => {
invocation.arg("delete").arg(key).arg(Vec::<u8>::new());
}
StorageOperation::Expire { key, ttl_seconds } => {
invocation
.arg("expire")
.arg(key)
.arg(ttl_seconds.to_string());
}
StorageOperation::RangeAdd { key, value } => {
invocation.arg("range_add").arg(key).arg(value);
}
StorageOperation::RangeRemove { key, value } => {
invocation.arg("range_remove").arg(key).arg(value);
}
StorageOperation::RangeExpire {
key,
value,
ttl_seconds,
} => {
invocation
.arg("range_expire")
.arg(key)
.arg(value)
.arg(ttl_seconds.to_string());
}
StorageOperation::RangeClear { key } => {
invocation.arg("range_clear").arg(key).arg(Vec::<u8>::new());
}
}
}
let mut connection = self.connection.clone();
invocation.invoke_async::<_, i32>(&mut connection).await?;
Ok(())
}
async fn range_page(
&self,
key: &str,
cursor: Option<i64>,
limit: usize,
order: RangeOrder,
) -> anyhow::Result<RangePage> {
self.remove_expired(key).await?;
let mut command = match order {
RangeOrder::OldestFirst => {
let minimum =
cursor.map_or_else(|| "-inf".to_string(), |value| format!("({value}"));
let mut command = redis::cmd("ZRANGEBYSCORE");
command.arg(key).arg(minimum).arg("+inf");
command
}
RangeOrder::NewestFirst => {
let maximum =
cursor.map_or_else(|| "+inf".to_string(), |value| format!("({value}"));
let mut command = redis::cmd("ZREVRANGEBYSCORE");
command.arg(key).arg(maximum).arg("-inf");
command
}
};
command.arg("WITHSCORES").arg("LIMIT").arg(0).arg(limit);
let mut connection = self.connection.clone();
let values = command
.query_async::<_, Vec<(Vec<u8>, f64)>>(&mut connection)
.await?;
let items = values
.into_iter()
.map(|(value, score)| {
Ok(RangeItem {
cursor: Self::score_to_cursor(score)?,
value,
})
})
.collect::<anyhow::Result<Vec<_>>>()?;
let next_cursor = (items.len() == limit)
.then(|| items.last().map(|item| item.cursor))
.flatten();
Ok(RangePage { items, next_cursor })
}
async fn range_count(&self, key: &str) -> anyhow::Result<usize> {
self.remove_expired(key).await?;
let mut connection = self.connection.clone();
Ok(connection.zcard(key).await?)
}
async fn job_index_upsert(&self, namespace: &str, mut row: JobIndexRow) -> anyhow::Result<()> {
let mut connection = self.connection.clone();
let row_key = format!("{namespace}:jobidx:row:{}", row.job_id);
let existing: Option<Vec<u8>> = connection.get(&row_key).await?;
if let Some(existing) = existing
.as_deref()
.map(crate::encoder::decode::<JobIndexRow>)
{
let existing = existing?;
if row.revision < existing.revision {
return Ok(());
}
row.created_at = existing.created_at;
if row.wait_ms.is_none() {
row.wait_ms = existing.wait_ms;
}
if row.wait_mode.is_none() {
row.wait_mode = existing.wait_mode.clone();
}
if row.parent_job_id.is_none() {
row.parent_job_id = existing.parent_job_id.clone();
}
if existing.stage != row.stage {
let old_stage_key = format!("{namespace}:jobidx:stage:{}", existing.stage);
let _: () = connection.zrem(&old_stage_key, &row.job_id).await?;
}
if existing.parent_job_id != row.parent_job_id {
if let Some(parent) = &existing.parent_job_id {
let old_parent_key = format!("{namespace}:jobidx:parent:{parent}");
let _: () = connection.zrem(&old_parent_key, &row.job_id).await?;
}
}
}
connection
.set::<_, _, ()>(&row_key, crate::encoder::encode(&row)?)
.await?;
let stage_key = format!("{namespace}:jobidx:stage:{}", row.stage);
connection
.zadd::<_, _, _, ()>(&stage_key, &row.job_id, row.stage_date.timestamp_millis())
.await?;
if let (Some(topic), Some(partition), Some(sequence)) =
(&row.topic, row.partition, row.sequence)
{
let partition_key = format!("{namespace}:jobidx:partition:{topic}:{partition}");
connection
.zadd::<_, _, _, ()>(&partition_key, &row.job_id, sequence)
.await?;
}
if let Some(parent) = &row.parent_job_id {
let parent_key = format!("{namespace}:jobidx:parent:{parent}");
connection
.zadd::<_, _, _, ()>(&parent_key, &row.job_id, row.stage_date.timestamp_millis())
.await?;
}
Ok(())
}
async fn job_index_stage_counts(
&self,
namespace: &str,
) -> anyhow::Result<std::collections::HashMap<String, usize>> {
let mut connection = self.connection.clone();
let pattern = format!("{namespace}:jobidx:stage:*");
let keys: Vec<String> = connection.keys(&pattern).await?;
let prefix = format!("{namespace}:jobidx:stage:");
let mut counts = std::collections::HashMap::new();
for key in keys {
let Some(stage) = key.strip_prefix(&prefix) else {
continue;
};
let count: usize = connection.zcard(&key).await?;
if count > 0 {
counts.insert(stage.to_string(), count);
}
}
Ok(counts)
}
async fn job_index_list_by_stage(
&self,
namespace: &str,
stage: &str,
cursor: Option<i64>,
limit: usize,
) -> anyhow::Result<JobIndexPage> {
let mut connection = self.connection.clone();
let key = format!("{namespace}:jobidx:stage:{stage}");
let max = cursor.map(|c| c - 1).unwrap_or(i64::MAX);
let job_ids: Vec<String> = connection
.zrevrangebyscore_limit(&key, max, i64::MIN, 0, isize::try_from(limit)?)
.await?;
let items = self.job_index_hydrate(namespace, &job_ids).await?;
let next_cursor = (job_ids.len() == limit)
.then(|| items.last().map(|item| item.stage_date.timestamp_millis()))
.flatten();
Ok(JobIndexPage { items, next_cursor })
}
async fn job_index_list_by_partition(
&self,
namespace: &str,
topic: &str,
partition: u32,
cursor: Option<i64>,
limit: usize,
) -> anyhow::Result<JobIndexPage> {
let mut connection = self.connection.clone();
let key = format!("{namespace}:jobidx:partition:{topic}:{partition}");
let max = cursor.map(|c| c - 1).unwrap_or(i64::MAX);
let job_ids: Vec<String> = connection
.zrevrangebyscore_limit(&key, max, i64::MIN, 0, isize::try_from(limit)?)
.await?;
let items = self.job_index_hydrate(namespace, &job_ids).await?;
let next_cursor = (job_ids.len() == limit)
.then(|| items.last().and_then(|item| item.sequence))
.flatten();
Ok(JobIndexPage { items, next_cursor })
}
async fn job_index_partition_neighbor(
&self,
namespace: &str,
topic: &str,
partition: u32,
sequence: i64,
older: bool,
) -> anyhow::Result<Option<JobIndexRow>> {
let mut connection = self.connection.clone();
let key = format!("{namespace}:jobidx:partition:{topic}:{partition}");
let target = if older { sequence - 1 } else { sequence + 1 };
let job_ids: Vec<String> = connection.zrangebyscore(&key, target, target).await?;
let Some(job_id) = job_ids.into_iter().next() else {
return Ok(None);
};
self.job_index_get(namespace, &job_id).await
}
async fn job_index_list_continuations(
&self,
namespace: &str,
parent_job_id: &str,
limit: usize,
) -> anyhow::Result<Vec<JobIndexRow>> {
let mut connection = self.connection.clone();
let key = format!("{namespace}:jobidx:parent:{parent_job_id}");
let job_ids: Vec<String> = connection
.zrangebyscore_limit(&key, i64::MIN, i64::MAX, 0, isize::try_from(limit)?)
.await?;
self.job_index_hydrate(namespace, &job_ids).await
}
async fn job_index_truncate_partition(
&self,
namespace: &str,
topic: &str,
partition: u32,
keep_last: usize,
) -> anyhow::Result<usize> {
let mut connection = self.connection.clone();
let key = format!("{namespace}:jobidx:partition:{topic}:{partition}");
let total: usize = connection.zcard(&key).await?;
let Some(to_remove) = total.checked_sub(keep_last) else {
return Ok(0);
};
if to_remove == 0 {
return Ok(0);
}
let job_ids: Vec<String> = connection
.zrange(&key, 0, isize::try_from(to_remove)? - 1)
.await?;
for job_id in &job_ids {
let _: () = connection.zrem(&key, job_id).await?;
let row_key = format!("{namespace}:jobidx:row:{job_id}");
let _: () = connection.del(&row_key).await?;
}
Ok(job_ids.len())
}
async fn job_index_record_transition(
&self,
namespace: &str,
stage: &str,
at: UtcDateTime,
) -> anyhow::Result<()> {
let mut connection = self.connection.clone();
let key = format!("{namespace}:jobidx:txn:{stage}:{}", minute_bucket(at));
let _: i64 = connection.incr(&key, 1).await?;
let _: bool = connection.expire(&key, 180).await?;
Ok(())
}
async fn job_index_record_metrics_batch(
&self,
namespace: &str,
batch: &crate::storage::JobIndexMetricsBatch,
) -> anyhow::Result<()> {
let mut connection = self.connection.clone();
for (stage, at, count) in &batch.transitions {
let key = format!("{namespace}:jobidx:txn:{stage}:{}", minute_bucket(*at));
let _: i64 = connection.incr(&key, *count).await?;
let _: bool = connection.expire(&key, 180).await?;
}
for (mode, at, count, sum_ms, max_ms) in &batch.waits {
let key = format!("{namespace}:jobidx:wait:{mode}:{}", minute_bucket(*at));
let _: i64 = connection.hincr(&key, "count", *count).await?;
let _: i64 = connection.hincr(&key, "sum_ms", *sum_ms).await?;
let current_max: Option<i64> = connection.hget(&key, "max_ms").await?;
if current_max.unwrap_or(0) < *max_ms {
let _: () = connection.hset(&key, "max_ms", *max_ms).await?;
}
let _: bool = connection.expire(&key, 180).await?;
}
Ok(())
}
async fn job_index_recent_transition_count(
&self,
namespace: &str,
stage: &str,
now: UtcDateTime,
) -> anyhow::Result<usize> {
let mut connection = self.connection.clone();
let bucket = minute_bucket(now);
let mut total = 0usize;
for b in [bucket - 1, bucket] {
let key = format!("{namespace}:jobidx:txn:{stage}:{b}");
let count: Option<usize> = connection.get(&key).await?;
total += count.unwrap_or(0);
}
Ok(total)
}
async fn job_index_record_wait(
&self,
namespace: &str,
mode: &str,
wait_ms: i64,
at: UtcDateTime,
) -> anyhow::Result<()> {
let mut connection = self.connection.clone();
let key = format!("{namespace}:jobidx:wait:{mode}:{}", minute_bucket(at));
let _: i64 = connection.hincr(&key, "count", 1).await?;
let _: i64 = connection.hincr(&key, "sum_ms", wait_ms).await?;
let current_max: Option<i64> = connection.hget(&key, "max_ms").await?;
if current_max.unwrap_or(0) < wait_ms {
let _: () = connection.hset(&key, "max_ms", wait_ms).await?;
}
let _: bool = connection.expire(&key, 180).await?;
Ok(())
}
async fn job_index_recent_wait_stats(
&self,
namespace: &str,
mode: &str,
now: UtcDateTime,
) -> anyhow::Result<JobIndexWaitStats> {
let mut connection = self.connection.clone();
let bucket = minute_bucket(now);
let (mut count, mut sum_ms, mut max_ms) = (0i64, 0i64, 0i64);
for b in [bucket - 1, bucket] {
let key = format!("{namespace}:jobidx:wait:{mode}:{b}");
let values: std::collections::HashMap<String, i64> = connection.hgetall(&key).await?;
count += values.get("count").copied().unwrap_or(0);
sum_ms += values.get("sum_ms").copied().unwrap_or(0);
max_ms = max_ms.max(values.get("max_ms").copied().unwrap_or(0));
}
let count = usize::try_from(count)?;
Ok(JobIndexWaitStats {
count,
avg_ms: if count == 0 {
0
} else {
sum_ms / i64::try_from(count)?
},
max_ms,
})
}
async fn job_index_sweep_expired(
&self,
_namespace: &str,
_now: UtcDateTime,
_limit: usize,
) -> anyhow::Result<usize> {
Ok(0)
}
}
impl Redis {
async fn job_index_get(
&self,
namespace: &str,
job_id: &str,
) -> anyhow::Result<Option<JobIndexRow>> {
let mut connection = self.connection.clone();
let row_key = format!("{namespace}:jobidx:row:{job_id}");
let value: Option<Vec<u8>> = connection.get(&row_key).await?;
value
.map(|value| crate::encoder::decode::<JobIndexRow>(&value))
.transpose()
}
async fn job_index_hydrate(
&self,
namespace: &str,
job_ids: &[String],
) -> anyhow::Result<Vec<JobIndexRow>> {
let now = chrono::Utc::now();
let mut items = Vec::with_capacity(job_ids.len());
for job_id in job_ids {
if let Some(row) = self.job_index_get(namespace, job_id).await? {
if row.date_expire.is_some_and(|expire| expire <= now) {
continue;
}
items.push(row);
}
}
Ok(items)
}
}