fix(runtime-state): govern redis connections

This commit is contained in:
fawney19
2026-05-21 22:51:57 +08:00
parent b8a65cbdec
commit ab0a90de97
24 changed files with 3155 additions and 1064 deletions

File diff suppressed because it is too large Load Diff

View File

@@ -4,7 +4,7 @@ use std::time::{Duration, Instant};
use tokio::sync::Mutex;
use crate::{RuntimeQueueEntry, RuntimeQueueReclaimConfig};
use crate::{DataLayerError, RuntimeQueueEntry, RuntimeQueueReclaimConfig};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemoryRuntimeStateConfig {
@@ -37,9 +37,9 @@ pub(crate) struct MemoryRuntimeBackend {
config: MemoryRuntimeStateConfig,
kv: Mutex<HashMap<String, MemoryKvEntry>>,
counters: Mutex<HashMap<String, MemoryCounterEntry>>,
sets: Mutex<HashMap<String, BTreeSet<String>>>,
scores: Mutex<HashMap<String, BTreeMap<String, f64>>>,
queues: Mutex<HashMap<String, VecDeque<RuntimeQueueEntry>>>,
sets: Mutex<HashMap<String, MemorySetEntry>>,
scores: Mutex<HashMap<String, MemoryScoreEntry>>,
queues: Mutex<HashMap<String, MemoryQueueStream>>,
queue_seq: AtomicU64,
locks: Mutex<HashMap<String, MemoryLockEntry>>,
semaphores: Mutex<HashMap<String, BTreeMap<String, u64>>>,
@@ -52,6 +52,80 @@ struct MemoryCounterEntry {
expires_at: Instant,
}
#[derive(Debug, Default)]
struct MemorySetEntry {
members: BTreeSet<String>,
expires_at: Option<Instant>,
}
#[derive(Debug, Default)]
struct MemoryScoreEntry {
scores: BTreeMap<String, f64>,
expires_at: Option<Instant>,
}
#[derive(Debug, Default)]
struct MemoryQueueStream {
entries: VecDeque<MemoryQueuedEntry>,
groups: HashMap<String, MemoryConsumerGroup>,
expires_at: Option<Instant>,
}
trait MemoryExpiringKey {
fn is_expired(&self, now: Instant) -> bool;
fn set_expires_at(&mut self, expires_at: Instant);
}
impl MemoryExpiringKey for MemorySetEntry {
fn is_expired(&self, now: Instant) -> bool {
self.expires_at.is_some_and(|expires_at| now >= expires_at)
}
fn set_expires_at(&mut self, expires_at: Instant) {
self.expires_at = Some(expires_at);
}
}
impl MemoryExpiringKey for MemoryScoreEntry {
fn is_expired(&self, now: Instant) -> bool {
self.expires_at.is_some_and(|expires_at| now >= expires_at)
}
fn set_expires_at(&mut self, expires_at: Instant) {
self.expires_at = Some(expires_at);
}
}
impl MemoryExpiringKey for MemoryQueueStream {
fn is_expired(&self, now: Instant) -> bool {
self.expires_at.is_some_and(|expires_at| now >= expires_at)
}
fn set_expires_at(&mut self, expires_at: Instant) {
self.expires_at = Some(expires_at);
}
}
#[derive(Debug, Clone)]
struct MemoryQueuedEntry {
sequence: u64,
entry: RuntimeQueueEntry,
}
#[derive(Debug, Default)]
struct MemoryConsumerGroup {
last_delivered_sequence: u64,
pending: BTreeMap<String, MemoryPendingQueueEntry>,
}
#[derive(Debug, Clone)]
struct MemoryPendingQueueEntry {
sequence: u64,
entry: RuntimeQueueEntry,
consumer: String,
delivered_at: Instant,
}
#[derive(Debug, Clone)]
pub(crate) struct MemoryLockEntry {
pub(crate) token: String,
@@ -143,16 +217,66 @@ impl MemoryRuntimeBackend {
}
pub(crate) async fn kv_delete(&self, key: &str) -> bool {
self.kv.lock().await.remove(key).is_some()
let kv_deleted = self.kv.lock().await.remove(key).is_some();
let set_deleted = self.sets.lock().await.remove(key).is_some();
let score_deleted = self.scores.lock().await.remove(key).is_some();
let queue_deleted = self.queues.lock().await.remove(key).is_some();
kv_deleted || set_deleted || score_deleted || queue_deleted
}
pub(crate) async fn kv_delete_many(&self, keys: &[String]) -> usize {
let keys = keys.iter().cloned().collect::<BTreeSet<_>>();
let mut deleted = BTreeSet::new();
let mut kv = self.kv.lock().await;
keys.iter().filter(|key| kv.remove(*key).is_some()).count()
for key in &keys {
if kv.remove(key).is_some() {
deleted.insert(key.clone());
}
}
drop(kv);
let mut sets = self.sets.lock().await;
for key in &keys {
if sets.remove(key).is_some() {
deleted.insert(key.clone());
}
}
drop(sets);
let mut scores = self.scores.lock().await;
for key in &keys {
if scores.remove(key).is_some() {
deleted.insert(key.clone());
}
}
drop(scores);
let mut queues = self.queues.lock().await;
for key in &keys {
if queues.remove(key).is_some() {
deleted.insert(key.clone());
}
}
deleted.len()
}
pub(crate) async fn kv_exists(&self, key: &str) -> bool {
self.kv_get(key).await.is_some()
if self.kv_get(key).await.is_some() {
return true;
}
let now = Instant::now();
let mut sets = self.sets.lock().await;
prune_memory_key(&mut sets, key, now);
if sets.contains_key(key) {
return true;
}
drop(sets);
let mut scores = self.scores.lock().await;
prune_memory_key(&mut scores, key, now);
if scores.contains_key(key) {
return true;
}
drop(scores);
let mut queues = self.queues.lock().await;
prune_memory_key(&mut queues, key, now);
queues.contains_key(key)
}
pub(crate) async fn kv_ttl_seconds(&self, key: &str) -> Option<i64> {
@@ -177,16 +301,77 @@ impl MemoryRuntimeBackend {
)
}
pub(crate) async fn key_expire(&self, key: &str, ttl: Duration) -> bool {
let now = Instant::now();
if ttl.is_zero() {
let kv_deleted = self.kv.lock().await.remove(key).is_some();
let set_deleted = self.sets.lock().await.remove(key).is_some();
let score_deleted = self.scores.lock().await.remove(key).is_some();
let queue_deleted = self.queues.lock().await.remove(key).is_some();
return kv_deleted || set_deleted || score_deleted || queue_deleted;
}
let expires_at = now + ttl;
{
let mut kv = self.kv.lock().await;
if let Some(entry) = kv.get_mut(key) {
if entry.is_expired(now) {
kv.remove(key);
} else {
entry.expires_at = Some(expires_at);
return true;
}
}
}
if set_memory_key_expiry(&self.sets, key, expires_at, now).await {
return true;
}
if set_memory_key_expiry(&self.scores, key, expires_at, now).await {
return true;
}
if set_memory_key_expiry(&self.queues, key, expires_at, now).await {
return true;
}
false
}
pub(crate) async fn kv_scan(&self, pattern: &str) -> Vec<String> {
let now = Instant::now();
let mut keys = BTreeSet::new();
let mut kv = self.kv.lock().await;
prune_kv(&mut kv, Instant::now());
let mut keys = kv
.keys()
.filter(|key| key_matches_pattern(key, pattern))
.cloned()
.collect::<Vec<_>>();
keys.sort();
keys
prune_kv(&mut kv, now);
keys.extend(
kv.keys()
.filter(|key| key_matches_pattern(key, pattern))
.cloned(),
);
drop(kv);
let mut sets = self.sets.lock().await;
prune_expiring_map(&mut sets, now);
keys.extend(
sets.keys()
.filter(|key| key_matches_pattern(key, pattern))
.cloned(),
);
drop(sets);
let mut scores = self.scores.lock().await;
prune_expiring_map(&mut scores, now);
keys.extend(
scores
.keys()
.filter(|key| key_matches_pattern(key, pattern))
.cloned(),
);
drop(scores);
let mut queues = self.queues.lock().await;
prune_expiring_map(&mut queues, now);
keys.extend(
queues
.keys()
.filter(|key| key_matches_pattern(key, pattern))
.cloned(),
);
keys.into_iter().collect()
}
pub(crate) async fn check_and_consume_rate_limit(
@@ -231,6 +416,8 @@ impl MemoryRuntimeBackend {
}
let mut remaining = None::<u32>;
let mut user_next = None::<u32>;
let mut key_next = None::<u32>;
let expires_at = now + ttl;
if user_limit > 0 {
let next = counters
@@ -247,6 +434,7 @@ impl MemoryRuntimeBackend {
})
.value;
remaining = Some(user_limit.saturating_sub(next));
user_next = Some(next);
}
if key_limit > 0 {
let next = counters
@@ -264,6 +452,15 @@ impl MemoryRuntimeBackend {
.value;
let key_remaining = key_limit.saturating_sub(next);
remaining = Some(remaining.map_or(key_remaining, |value| value.min(key_remaining)));
key_next = Some(next);
}
drop(counters);
if let Some(next) = user_next {
self.kv_set(user_key, next.to_string(), Some(ttl)).await;
}
if let Some(next) = key_next {
self.kv_set(key_key, next.to_string(), Some(ttl)).await;
}
Ok(crate::RateLimitCheck::Allowed {
remaining: remaining.unwrap_or(0),
@@ -271,11 +468,11 @@ impl MemoryRuntimeBackend {
}
pub(crate) async fn set_add(&self, key: &str, member: &str) -> bool {
self.sets
.lock()
.await
.entry(key.to_string())
let mut sets = self.sets.lock().await;
prune_memory_key(&mut sets, key, Instant::now());
sets.entry(key.to_string())
.or_default()
.members
.insert(member.to_string())
}
@@ -283,82 +480,113 @@ impl MemoryRuntimeBackend {
let Ok(mut sets) = self.sets.try_lock() else {
return false;
};
prune_memory_key(&mut sets, key, Instant::now());
sets.entry(key.to_string())
.or_default()
.members
.insert(member.to_string())
}
pub(crate) async fn set_remove(&self, key: &str, member: &str) -> bool {
self.sets
.lock()
.await
.get_mut(key)
.is_some_and(|set| set.remove(member))
let mut sets = self.sets.lock().await;
prune_memory_key(&mut sets, key, Instant::now());
sets.get_mut(key)
.is_some_and(|entry| entry.members.remove(member))
}
pub(crate) async fn set_members(&self, key: &str) -> Vec<String> {
self.sets
.lock()
.await
.get(key)
.map(|set| set.iter().cloned().collect())
let mut sets = self.sets.lock().await;
prune_memory_key(&mut sets, key, Instant::now());
sets.get(key)
.map(|entry| entry.members.iter().cloned().collect())
.unwrap_or_default()
}
pub(crate) async fn set_len(&self, key: &str) -> usize {
self.sets.lock().await.get(key).map_or(0, BTreeSet::len)
let mut sets = self.sets.lock().await;
prune_memory_key(&mut sets, key, Instant::now());
sets.get(key).map_or(0, |entry| entry.members.len())
}
pub(crate) async fn score_set(&self, key: &str, member: &str, score: f64) {
self.scores
.lock()
.await
let mut scores = self.scores.lock().await;
prune_memory_key(&mut scores, key, Instant::now());
scores
.entry(key.to_string())
.or_default()
.scores
.insert(member.to_string(), score);
}
pub(crate) async fn score_many(&self, key: &str, members: &[String]) -> Vec<Option<f64>> {
let scores = self.scores.lock().await;
let mut scores = self.scores.lock().await;
prune_memory_key(&mut scores, key, Instant::now());
members
.iter()
.map(|member| scores.get(key).and_then(|set| set.get(member)).copied())
.map(|member| {
scores
.get(key)
.and_then(|entry| entry.scores.get(member))
.copied()
})
.collect()
}
pub(crate) async fn score_range_by_min(&self, key: &str, min_score: f64) -> Vec<String> {
let scores = self.scores.lock().await;
let mut scores = self.scores.lock().await;
prune_memory_key(&mut scores, key, Instant::now());
scores
.get(key)
.map(|set| {
set.iter()
.filter(|(_, score)| **score >= min_score)
.map(|(member, _)| member.clone())
.collect()
})
.map(|entry| sorted_score_members(&entry.scores, |score| score >= min_score))
.unwrap_or_default()
}
pub(crate) async fn score_remove_by_score(&self, key: &str, max_score: f64) -> usize {
let mut scores = self.scores.lock().await;
let Some(set) = scores.get_mut(key) else {
prune_memory_key(&mut scores, key, Instant::now());
let Some(entry) = scores.get_mut(key) else {
return 0;
};
let before = set.len();
set.retain(|_, score| *score > max_score);
before.saturating_sub(set.len())
let before = entry.scores.len();
entry.scores.retain(|_, score| *score > max_score);
before.saturating_sub(entry.scores.len())
}
pub(crate) async fn score_remove(&self, key: &str, member: &str) -> bool {
self.scores
.lock()
.await
let mut scores = self.scores.lock().await;
prune_memory_key(&mut scores, key, Instant::now());
scores
.get_mut(key)
.is_some_and(|set| set.remove(member).is_some())
.is_some_and(|entry| entry.scores.remove(member).is_some())
}
pub(crate) async fn score_remove_by_rank(&self, key: &str, start: i64, stop: i64) -> usize {
let mut scores = self.scores.lock().await;
prune_memory_key(&mut scores, key, Instant::now());
let Some(entry) = scores.get_mut(key) else {
return 0;
};
let Some((start, stop)) = normalize_redis_rank_range(entry.scores.len(), start, stop)
else {
return 0;
};
let members = sorted_score_members(&entry.scores, |_| true);
let remove = members
.into_iter()
.enumerate()
.filter_map(|(index, member)| (index >= start && index <= stop).then_some(member))
.collect::<Vec<_>>();
let before = entry.scores.len();
for member in remove {
entry.scores.remove(&member);
}
before.saturating_sub(entry.scores.len())
}
pub(crate) async fn score_len(&self, key: &str) -> usize {
self.scores.lock().await.get(key).map_or(0, BTreeMap::len)
let mut scores = self.scores.lock().await;
prune_memory_key(&mut scores, key, Instant::now());
scores.get(key).map_or(0, |entry| entry.scores.len())
}
pub(crate) async fn queue_append(
@@ -367,51 +595,206 @@ impl MemoryRuntimeBackend {
fields: BTreeMap<String, String>,
maxlen: Option<usize>,
) -> String {
let id = format!(
"{}-0",
self.queue_seq
.fetch_add(1, Ordering::Relaxed)
.saturating_add(1)
);
let sequence = self
.queue_seq
.fetch_add(1, Ordering::Relaxed)
.saturating_add(1);
let id = format!("{sequence}-0");
let mut queues = self.queues.lock().await;
let queue = queues.entry(stream.to_string()).or_default();
queue.push_back(RuntimeQueueEntry {
id: id.clone(),
fields,
prune_memory_key(&mut queues, stream, Instant::now());
let stream_state = queues.entry(stream.to_string()).or_default();
stream_state.entries.push_back(MemoryQueuedEntry {
sequence,
entry: RuntimeQueueEntry {
id: id.clone(),
fields,
},
});
if let Some(maxlen) = maxlen.filter(|value| *value > 0) {
while queue.len() > maxlen {
queue.pop_front();
while stream_state.entries.len() > maxlen {
let Some(removed) = stream_state.entries.pop_front() else {
break;
};
remove_pending_from_all_groups(stream_state, &removed.entry.id);
}
}
id
}
pub(crate) async fn queue_read(&self, stream: &str, count: usize) -> Vec<RuntimeQueueEntry> {
pub(crate) async fn queue_ensure_consumer_group(
&self,
stream: &str,
group: &str,
start_id: &str,
) -> Result<(), DataLayerError> {
let mut queues = self.queues.lock().await;
let Some(queue) = queues.get_mut(stream) else {
return Vec::new();
};
let mut entries = Vec::new();
for _ in 0..count.max(1) {
let Some(entry) = queue.pop_front() else {
break;
};
entries.push(entry);
prune_memory_key(&mut queues, stream, Instant::now());
let stream_state = queues.entry(stream.to_string()).or_default();
if stream_state.groups.contains_key(group) {
return Ok(());
}
let last_delivered_sequence = match start_id {
"$" => stream_state
.entries
.back()
.map(|entry| entry.sequence)
.unwrap_or_default(),
_ => parse_memory_stream_sequence(start_id)?,
};
stream_state.groups.insert(
group.to_string(),
MemoryConsumerGroup {
last_delivered_sequence,
pending: BTreeMap::new(),
},
);
Ok(())
}
pub(crate) async fn queue_read(
&self,
stream: &str,
group: &str,
consumer: &str,
count: usize,
block_ms: Option<u64>,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
let deadline = block_ms.map(|value| Instant::now() + Duration::from_millis(value.max(1)));
loop {
let entries = {
let mut queues = self.queues.lock().await;
prune_memory_key(&mut queues, stream, Instant::now());
let Some(stream_state) = queues.get_mut(stream) else {
return Err(DataLayerError::InvalidInput(format!(
"runtime queue stream {stream} does not exist"
)));
};
let Some(group_state) = stream_state.groups.get_mut(group) else {
return Err(DataLayerError::InvalidInput(format!(
"runtime queue group {group} does not exist for stream {stream}"
)));
};
let now = Instant::now();
let mut delivered = Vec::new();
let last_delivered_sequence = group_state.last_delivered_sequence;
let queued_entries = stream_state
.entries
.iter()
.filter(|entry| entry.sequence > last_delivered_sequence)
.take(count.max(1))
.cloned()
.collect::<Vec<_>>();
for queued in queued_entries {
group_state.last_delivered_sequence = queued.sequence;
group_state.pending.insert(
queued.entry.id.clone(),
MemoryPendingQueueEntry {
sequence: queued.sequence,
entry: queued.entry.clone(),
consumer: consumer.to_string(),
delivered_at: now,
},
);
delivered.push(queued.entry.clone());
}
delivered
};
if !entries.is_empty() {
return Ok(entries);
}
let Some(deadline) = deadline else {
return Ok(Vec::new());
};
if Instant::now() >= deadline {
return Ok(Vec::new());
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
entries
}
pub(crate) async fn queue_claim_stale(
&self,
_stream: &str,
_config: RuntimeQueueReclaimConfig,
) -> Vec<RuntimeQueueEntry> {
Vec::new()
stream: &str,
group: &str,
consumer: &str,
start_id: &str,
config: RuntimeQueueReclaimConfig,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
let start_sequence = parse_memory_stream_sequence(start_id)?;
let min_idle = Duration::from_millis(config.min_idle_ms.max(1));
let now = Instant::now();
let mut queues = self.queues.lock().await;
prune_memory_key(&mut queues, stream, now);
let Some(stream_state) = queues.get_mut(stream) else {
return Err(DataLayerError::InvalidInput(format!(
"runtime queue stream {stream} does not exist"
)));
};
let Some(group_state) = stream_state.groups.get_mut(group) else {
return Err(DataLayerError::InvalidInput(format!(
"runtime queue group {group} does not exist for stream {stream}"
)));
};
let ids = group_state
.pending
.values()
.filter(|entry| entry.sequence >= start_sequence)
.filter(|entry| now.saturating_duration_since(entry.delivered_at) >= min_idle)
.map(|entry| (entry.sequence, entry.entry.id.clone()))
.collect::<Vec<_>>();
let mut ids = ids;
ids.sort_by_key(|(sequence, _)| *sequence);
let mut claimed = Vec::new();
for (_, id) in ids.into_iter().take(config.count.max(1)) {
if let Some(pending) = group_state.pending.get_mut(&id) {
pending.consumer = consumer.to_string();
pending.delivered_at = now;
claimed.push(pending.entry.clone());
}
}
Ok(claimed)
}
pub(crate) async fn queue_delete(&self, _stream: &str, _ids: &[String]) -> usize {
0
pub(crate) async fn queue_ack(
&self,
stream: &str,
group: &str,
ids: &[String],
) -> Result<usize, DataLayerError> {
let mut queues = self.queues.lock().await;
prune_memory_key(&mut queues, stream, Instant::now());
let Some(stream_state) = queues.get_mut(stream) else {
return Err(DataLayerError::InvalidInput(format!(
"runtime queue stream {stream} does not exist"
)));
};
let Some(group_state) = stream_state.groups.get_mut(group) else {
return Err(DataLayerError::InvalidInput(format!(
"runtime queue group {group} does not exist for stream {stream}"
)));
};
Ok(ids
.iter()
.filter(|id| group_state.pending.remove(*id).is_some())
.count())
}
pub(crate) async fn queue_delete(&self, stream: &str, ids: &[String]) -> usize {
let mut queues = self.queues.lock().await;
prune_memory_key(&mut queues, stream, Instant::now());
let Some(stream_state) = queues.get_mut(stream) else {
return 0;
};
let ids = ids.iter().cloned().collect::<BTreeSet<_>>();
let before = stream_state.entries.len();
stream_state
.entries
.retain(|entry| !ids.contains(&entry.entry.id));
for id in &ids {
remove_pending_from_all_groups(stream_state, id);
}
before.saturating_sub(stream_state.entries.len())
}
pub(crate) async fn lock_try_acquire(
@@ -530,6 +913,43 @@ fn prune_kv(kv: &mut HashMap<String, MemoryKvEntry>, now: Instant) {
kv.retain(|_, entry| !entry.is_expired(now));
}
fn prune_memory_key<T>(values: &mut HashMap<String, T>, key: &str, now: Instant)
where
T: MemoryExpiringKey,
{
if values.get(key).is_some_and(|entry| entry.is_expired(now)) {
values.remove(key);
}
}
fn prune_expiring_map<T>(values: &mut HashMap<String, T>, now: Instant)
where
T: MemoryExpiringKey,
{
values.retain(|_, entry| !entry.is_expired(now));
}
async fn set_memory_key_expiry<T>(
values: &Mutex<HashMap<String, T>>,
key: &str,
expires_at: Instant,
now: Instant,
) -> bool
where
T: MemoryExpiringKey,
{
let mut values = values.lock().await;
if values.get(key).is_some_and(|entry| entry.is_expired(now)) {
values.remove(key);
return false;
}
let Some(entry) = values.get_mut(key) else {
return false;
};
entry.set_expires_at(expires_at);
true
}
pub(crate) fn key_matches_pattern(key: &str, pattern: &str) -> bool {
match pattern.strip_suffix('*') {
Some(prefix) => key.starts_with(prefix),
@@ -537,6 +957,58 @@ pub(crate) fn key_matches_pattern(key: &str, pattern: &str) -> bool {
}
}
fn sorted_score_members<F>(scores: &BTreeMap<String, f64>, include: F) -> Vec<String>
where
F: Fn(f64) -> bool,
{
let mut entries = scores
.iter()
.filter_map(|(member, score)| include(*score).then_some((member.clone(), *score)))
.collect::<Vec<_>>();
entries.sort_by(|(left_member, left_score), (right_member, right_score)| {
left_score
.total_cmp(right_score)
.then_with(|| left_member.cmp(right_member))
});
entries.into_iter().map(|(member, _)| member).collect()
}
fn normalize_redis_rank_range(len: usize, start: i64, stop: i64) -> Option<(usize, usize)> {
if len == 0 {
return None;
}
let len = i64::try_from(len).ok()?;
let mut start = if start < 0 { len + start } else { start };
let mut stop = if stop < 0 { len + stop } else { stop };
if start < 0 {
start = 0;
}
if stop < 0 || start >= len || start > stop {
return None;
}
if stop >= len {
stop = len - 1;
}
Some((usize::try_from(start).ok()?, usize::try_from(stop).ok()?))
}
fn remove_pending_from_all_groups(stream: &mut MemoryQueueStream, id: &str) {
for group in stream.groups.values_mut() {
group.pending.remove(id);
}
}
fn parse_memory_stream_sequence(id: &str) -> Result<u64, DataLayerError> {
let Some((sequence, _)) = id.split_once('-') else {
return Err(DataLayerError::InvalidInput(format!(
"runtime queue stream id {id} must use redis stream id format"
)));
};
sequence.parse::<u64>().map_err(|err| {
DataLayerError::InvalidInput(format!("runtime queue stream id {id} is invalid: {err}"))
})
}
fn unix_time_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)

View File

@@ -1,8 +1,13 @@
use crate::error::RedisResultExt;
use crate::redis::RedisKeyspace;
use crate::DataLayerError;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tracing::info;
pub type RedisClient = redis::Client;
pub(crate) type RedisClient = redis::Client;
pub(crate) type RedisManagedConnection = redis::aio::ConnectionManager;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub struct RedisClientConfig {
@@ -30,23 +35,209 @@ impl RedisClientConfig {
}
#[derive(Debug, Clone)]
pub struct RedisClientFactory {
pub(crate) struct RedisClientFactory {
config: RedisClientConfig,
}
impl RedisClientFactory {
pub fn new(config: RedisClientConfig) -> Result<Self, DataLayerError> {
pub(crate) fn new(config: RedisClientConfig) -> Result<Self, DataLayerError> {
config.validate()?;
Ok(Self { config })
}
pub fn config(&self) -> &RedisClientConfig {
pub(crate) fn config(&self) -> &RedisClientConfig {
&self.config
}
pub fn connect_lazy(&self) -> Result<RedisClient, DataLayerError> {
pub(crate) fn connect_lazy(&self) -> Result<RedisClient, DataLayerError> {
RedisClient::open(self.config.url.clone()).map_redis_err()
}
pub(crate) async fn connect_router(
&self,
command_timeout_ms: Option<u64>,
) -> Result<RedisConnectionRouter, DataLayerError> {
RedisConnectionRouter::connect(self.connect_lazy()?, command_timeout_ms).await
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum RedisConnectionLane {
Fast,
Stream,
BlockingStream,
Admin,
}
impl RedisConnectionLane {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Fast => "fast",
Self::Stream => "stream",
Self::BlockingStream => "blocking_stream",
Self::Admin => "admin",
}
}
}
#[derive(Clone)]
pub(crate) struct RedisConnectionRouter {
fast: RedisManagedConnection,
stream: RedisManagedConnection,
blocking_stream: RedisManagedConnection,
admin: RedisManagedConnection,
metrics: Arc<RedisConnectionMetrics>,
}
impl std::fmt::Debug for RedisConnectionRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RedisConnectionRouter")
.field("lanes", &["fast", "stream", "blocking_stream", "admin"])
.finish()
}
}
impl RedisConnectionRouter {
pub(crate) async fn connect(
client: RedisClient,
command_timeout_ms: Option<u64>,
) -> Result<Self, DataLayerError> {
let fast = connect_lane(
&client,
connection_manager_config(command_timeout_ms),
RedisConnectionLane::Fast,
)
.await?;
let stream = connect_lane(
&client,
connection_manager_config(command_timeout_ms),
RedisConnectionLane::Stream,
)
.await?;
let blocking_stream = connect_lane(
&client,
connection_manager_config(command_timeout_ms),
RedisConnectionLane::BlockingStream,
)
.await?;
let admin = connect_lane(
&client,
connection_manager_config(command_timeout_ms),
RedisConnectionLane::Admin,
)
.await?;
info!(
redis_lanes = "fast,stream,blocking_stream,admin",
"runtime redis connection lanes initialized"
);
Ok(Self {
fast,
stream,
blocking_stream,
admin,
metrics: Arc::new(RedisConnectionMetrics::default()),
})
}
pub(crate) fn connection(&self, lane: RedisConnectionLane) -> RedisManagedConnection {
match lane {
RedisConnectionLane::Fast => self.fast.clone(),
RedisConnectionLane::Stream => self.stream.clone(),
RedisConnectionLane::BlockingStream => self.blocking_stream.clone(),
RedisConnectionLane::Admin => self.admin.clone(),
}
}
pub(crate) fn record_error(&self, lane: RedisConnectionLane) {
self.metrics
.for_lane(lane)
.errors
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_timeout(&self, lane: RedisConnectionLane) {
self.metrics
.for_lane(lane)
.timeouts
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn lane_diagnostics(&self) -> Vec<RedisLaneDiagnostics> {
[
RedisConnectionLane::Fast,
RedisConnectionLane::Stream,
RedisConnectionLane::BlockingStream,
RedisConnectionLane::Admin,
]
.into_iter()
.map(|lane| {
let metrics = self.metrics.for_lane(lane);
RedisLaneDiagnostics {
lane: lane.as_str(),
command_errors: metrics.errors.load(Ordering::Relaxed),
command_timeouts: metrics.timeouts.load(Ordering::Relaxed),
}
})
.collect()
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub struct RedisLaneDiagnostics {
pub lane: &'static str,
pub command_errors: u64,
pub command_timeouts: u64,
}
#[derive(Default)]
struct RedisConnectionMetrics {
fast: RedisLaneMetrics,
stream: RedisLaneMetrics,
blocking_stream: RedisLaneMetrics,
admin: RedisLaneMetrics,
}
impl RedisConnectionMetrics {
fn for_lane(&self, lane: RedisConnectionLane) -> &RedisLaneMetrics {
match lane {
RedisConnectionLane::Fast => &self.fast,
RedisConnectionLane::Stream => &self.stream,
RedisConnectionLane::BlockingStream => &self.blocking_stream,
RedisConnectionLane::Admin => &self.admin,
}
}
}
#[derive(Default)]
struct RedisLaneMetrics {
errors: AtomicU64,
timeouts: AtomicU64,
}
fn connection_manager_config(
command_timeout_ms: Option<u64>,
) -> redis::aio::ConnectionManagerConfig {
let mut config = redis::aio::ConnectionManagerConfig::new();
if let Some(timeout_ms) = command_timeout_ms {
config = config.set_connection_timeout(Duration::from_millis(timeout_ms));
}
config
}
async fn connect_lane(
client: &RedisClient,
config: redis::aio::ConnectionManagerConfig,
lane: RedisConnectionLane,
) -> Result<RedisManagedConnection, DataLayerError> {
client
.get_connection_manager_with_config(config)
.await
.map_err(|err| {
DataLayerError::Redis(format!(
"failed to initialize runtime redis {} lane: {err}",
lane.as_str()
))
})
}
#[cfg(test)]

View File

@@ -1,8 +1,10 @@
use std::future::Future;
use std::time::Duration;
use crate::error::RedisResultExt;
use crate::redis::{RedisClient, RedisKeyspace};
use crate::redis::{
run_lane_with_timeout, RedisClientConfig, RedisClientFactory, RedisConnectionLane,
RedisConnectionRouter, RedisKeyspace,
};
use crate::DataLayerError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -40,27 +42,35 @@ impl RedisKvRunnerConfig {
#[derive(Debug, Clone)]
pub struct RedisKvRunner {
client: RedisClient,
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
config: RedisKvRunnerConfig,
}
impl RedisKvRunner {
pub fn new(
client: RedisClient,
pub(crate) fn new(
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
config: RedisKvRunnerConfig,
) -> Result<Self, DataLayerError> {
config.validate()?;
Ok(Self {
client,
connections,
keyspace,
config,
})
}
pub fn client(&self) -> &RedisClient {
&self.client
pub async fn from_config(
config: RedisClientConfig,
runner_config: RedisKvRunnerConfig,
) -> Result<Self, DataLayerError> {
let factory = RedisClientFactory::new(config)?;
let keyspace = factory.config().keyspace();
let connections = factory
.connect_router(runner_config.command_timeout_ms)
.await?;
Self::new(connections, keyspace, runner_config)
}
pub fn keyspace(&self) -> &RedisKeyspace {
@@ -79,12 +89,8 @@ impl RedisKvRunner {
) -> Result<String, DataLayerError> {
let resolved_ttl = ttl_seconds.unwrap_or(self.config.default_ttl_seconds);
let namespaced_key = self.keyspace.key(key);
self.run_with_timeout("redis kv setex", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Fast, "redis kv setex", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
redis::cmd("SETEX")
.arg(&namespaced_key)
.arg(resolved_ttl)
@@ -98,12 +104,8 @@ impl RedisKvRunner {
pub async fn get(&self, key: &str) -> Result<Option<String>, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
self.run_with_timeout("redis kv get", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Fast, "redis kv get", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
redis::cmd("GET")
.arg(&namespaced_key)
.query_async(&mut connection)
@@ -115,12 +117,8 @@ impl RedisKvRunner {
pub async fn getdel(&self, key: &str) -> Result<Option<String>, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
self.run_with_timeout("redis kv getdel", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Fast, "redis kv getdel", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
redis::cmd("GETDEL")
.arg(&namespaced_key)
.query_async(&mut connection)
@@ -133,12 +131,8 @@ impl RedisKvRunner {
pub async fn exists(&self, key: &str) -> Result<bool, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
let exists = self
.run_with_timeout("redis kv exists", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
.run_with_timeout(RedisConnectionLane::Fast, "redis kv exists", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
redis::cmd("EXISTS")
.arg(&namespaced_key)
.query_async::<i64>(&mut connection)
@@ -151,12 +145,8 @@ impl RedisKvRunner {
pub async fn del(&self, key: &str) -> Result<i64, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
self.run_with_timeout("redis kv del", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Fast, "redis kv del", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
redis::cmd("DEL")
.arg(&namespaced_key)
.query_async(&mut connection)
@@ -168,49 +158,33 @@ impl RedisKvRunner {
async fn run_with_timeout<T, F>(
&self,
lane: RedisConnectionLane,
operation: &'static str,
future: F,
) -> Result<T, DataLayerError>
where
F: Future<Output = Result<T, DataLayerError>>,
{
if let Some(timeout_ms) = self.config.command_timeout_ms {
tokio::time::timeout(Duration::from_millis(timeout_ms), future)
.await
.map_err(|_| {
DataLayerError::TimedOut(format!("{operation} exceeded {timeout_ms}ms timeout"))
})?
} else {
future.await
}
run_lane_with_timeout(
&self.connections,
lane,
self.config.command_timeout_ms,
operation,
future,
)
.await
}
}
#[cfg(test)]
mod tests {
use super::{RedisKvRunner, RedisKvRunnerConfig};
use crate::redis::{RedisClientConfig, RedisClientFactory, RedisKeyspace};
fn build_runner() -> RedisKvRunner {
let config = RedisClientConfig {
url: "redis://localhost/0".to_string(),
key_prefix: Some("aether-test".to_string()),
};
let factory = RedisClientFactory::new(config).expect("redis factory");
let client = factory.connect_lazy().expect("connect");
let keyspace = factory.config().keyspace();
RedisKvRunner::new(client, keyspace, RedisKvRunnerConfig::default()).expect("runner build")
}
use super::RedisKvRunnerConfig;
#[test]
fn runner_reuses_client_keyspace_and_config() {
let runner = build_runner();
assert_eq!(
runner.keyspace().key("kv:setex:1"),
"aether-test:kv:setex:1"
);
assert_eq!(runner.config(), RedisKvRunnerConfig::default());
let _client = runner.client();
fn validates_default_config() {
RedisKvRunnerConfig::default()
.validate()
.expect("default kv config should be valid");
}
#[test]
@@ -219,17 +193,6 @@ mod tests {
command_timeout_ms: Some(100),
default_ttl_seconds: 0,
};
assert!(RedisKvRunner::new(
RedisClientFactory::new(RedisClientConfig {
url: "redis://localhost/0".to_string(),
key_prefix: Some("aether-test".to_string()),
})
.expect("redis factory")
.connect_lazy()
.expect("redis client"),
RedisKeyspace::new(Some("aether-test")),
config,
)
.is_err());
assert!(config.validate().is_err());
}
}

View File

@@ -1,8 +1,10 @@
use std::future::Future;
use std::time::Duration;
use crate::error::RedisResultExt;
use crate::redis::{RedisClient, RedisKeyspace};
use crate::redis::{
run_lane_with_timeout, RedisClientConfig, RedisClientFactory, RedisConnectionLane,
RedisConnectionRouter, RedisKeyspace,
};
use crate::DataLayerError;
use uuid::Uuid;
@@ -50,27 +52,35 @@ impl RedisLockRunnerConfig {
#[derive(Debug, Clone)]
pub struct RedisLockRunner {
client: RedisClient,
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
config: RedisLockRunnerConfig,
}
impl RedisLockRunner {
pub fn new(
client: RedisClient,
pub(crate) fn new(
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
config: RedisLockRunnerConfig,
) -> Result<Self, DataLayerError> {
config.validate()?;
Ok(Self {
client,
connections,
keyspace,
config,
})
}
pub fn client(&self) -> &RedisClient {
&self.client
pub async fn from_config(
config: RedisClientConfig,
runner_config: RedisLockRunnerConfig,
) -> Result<Self, DataLayerError> {
let factory = RedisClientFactory::new(config)?;
let keyspace = factory.config().keyspace();
let connections = factory
.connect_router(runner_config.command_timeout_ms)
.await?;
Self::new(connections, keyspace, runner_config)
}
pub fn keyspace(&self) -> &RedisKeyspace {
@@ -92,12 +102,8 @@ impl RedisLockRunner {
let ttl_ms = self.resolve_ttl_ms(ttl_ms)?;
let token = format!("{owner}:{}", Uuid::new_v4());
self.run_with_timeout("redis lock acquire", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Fast, "redis lock acquire", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
let status = redis::cmd("SET")
.arg(&key.0)
.arg(&token)
@@ -120,12 +126,8 @@ impl RedisLockRunner {
pub async fn release(&self, lease: &RedisLockLease) -> Result<bool, DataLayerError> {
validate_lease(lease)?;
self.run_with_timeout("redis lock release", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Fast, "redis lock release", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
let deleted = redis::Script::new(
"if redis.call('get', KEYS[1]) == ARGV[1] then \
return redis.call('del', KEYS[1]) \
@@ -151,12 +153,8 @@ impl RedisLockRunner {
validate_lease(lease)?;
let ttl_ms = self.resolve_ttl_ms(ttl_ms)?;
self.run_with_timeout("redis lock renew", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Fast, "redis lock renew", async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
let renewed = redis::Script::new(
"if redis.call('get', KEYS[1]) == ARGV[1] then \
return redis.call('pexpire', KEYS[1], ARGV[2]) \
@@ -177,21 +175,21 @@ impl RedisLockRunner {
async fn run_with_timeout<T, F>(
&self,
lane: RedisConnectionLane,
operation: &'static str,
future: F,
) -> Result<T, DataLayerError>
where
F: Future<Output = Result<T, DataLayerError>>,
{
if let Some(timeout_ms) = self.config.command_timeout_ms {
tokio::time::timeout(Duration::from_millis(timeout_ms), future)
.await
.map_err(|_| {
DataLayerError::TimedOut(format!("{operation} exceeded {timeout_ms}ms timeout"))
})?
} else {
future.await
}
run_lane_with_timeout(
&self.connections,
lane,
self.config.command_timeout_ms,
operation,
future,
)
.await
}
fn resolve_ttl_ms(&self, ttl_ms: Option<u64>) -> Result<u64, DataLayerError> {
@@ -241,29 +239,10 @@ fn validate_lease(lease: &RedisLockLease) -> Result<(), DataLayerError> {
#[cfg(test)]
mod tests {
use super::{RedisLockKey, RedisLockLease, RedisLockRunner, RedisLockRunnerConfig};
use crate::redis::{RedisClientConfig, RedisClientFactory};
fn sample_runner() -> RedisLockRunner {
let client = RedisClientFactory::new(RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
})
.expect("factory should build")
.connect_lazy()
.expect("client should build");
RedisLockRunner::new(
client,
RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
}
.keyspace(),
RedisLockRunnerConfig::default(),
)
.expect("runner should build")
}
use super::{
validate_key, validate_lease, validate_owner, RedisLockKey, RedisLockLease,
RedisLockRunnerConfig,
};
#[test]
fn validates_runner_config() {
@@ -283,41 +262,21 @@ mod tests {
#[test]
fn runner_reuses_client_and_keyspace() {
let runner = sample_runner();
assert_eq!(runner.config(), RedisLockRunnerConfig::default());
assert_eq!(runner.keyspace().lock_key("poller").0, "aether:lock:poller");
let _client_ref = runner.client();
RedisLockRunnerConfig::default()
.validate()
.expect("default lock config should be valid");
}
#[tokio::test]
async fn rejects_invalid_owner_or_lease_before_network() {
let runner = sample_runner();
assert!(runner
.try_acquire(&RedisLockKey("aether:lock:poller".to_string()), "", None)
.await
.is_err());
assert!(runner
.release(&RedisLockLease {
key: RedisLockKey("aether:lock:poller".to_string()),
owner: "worker-1".to_string(),
token: String::new(),
ttl_ms: 1_000,
})
.await
.is_err());
assert!(runner
.renew(
&RedisLockLease {
key: RedisLockKey("aether:lock:poller".to_string()),
owner: "worker-1".to_string(),
token: "token-1".to_string(),
ttl_ms: 1_000,
},
Some(0),
)
.await
.is_err());
#[test]
fn rejects_invalid_owner_or_lease_before_network() {
assert!(validate_owner("").is_err());
assert!(validate_key(&RedisLockKey(String::new())).is_err());
assert!(validate_lease(&RedisLockLease {
key: RedisLockKey("aether:lock:poller".to_string()),
owner: "worker-1".to_string(),
token: String::new(),
ttl_ms: 1_000,
})
.is_err());
}
}

View File

@@ -2,12 +2,14 @@ mod client;
mod kv;
mod lock;
mod namespace;
mod runtime;
mod stream;
pub use client::{RedisClient, RedisClientConfig, RedisClientFactory};
pub use client::{RedisClientConfig, RedisLaneDiagnostics};
pub use kv::{RedisKvRunner, RedisKvRunnerConfig};
pub use lock::{RedisLockKey, RedisLockLease, RedisLockRunner, RedisLockRunnerConfig};
pub use namespace::RedisKeyspace;
pub use runtime::RedisRuntimeDiagnostics;
pub use stream::{
RedisConsumerGroup, RedisConsumerName, RedisStreamEntry, RedisStreamName,
RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunner, RedisStreamRunnerConfig,
@@ -16,6 +18,9 @@ pub use stream::{
pub(crate) type RedisCmd = redis::Cmd;
pub(crate) type RedisScript = redis::Script;
pub(crate) use client::{RedisClientFactory, RedisConnectionLane, RedisConnectionRouter};
pub(crate) use runtime::RedisRuntimeRunner;
pub(crate) fn cmd(name: &str) -> RedisCmd {
redis::cmd(name)
}
@@ -23,3 +28,32 @@ pub(crate) fn cmd(name: &str) -> RedisCmd {
pub(crate) fn script(source: &str) -> RedisScript {
redis::Script::new(source)
}
pub(crate) async fn run_lane_with_timeout<T, F>(
connections: &RedisConnectionRouter,
lane: RedisConnectionLane,
timeout_ms: Option<u64>,
operation: &'static str,
future: F,
) -> Result<T, crate::DataLayerError>
where
F: std::future::Future<Output = Result<T, crate::DataLayerError>>,
{
let result = if let Some(timeout_ms) = timeout_ms {
match tokio::time::timeout(std::time::Duration::from_millis(timeout_ms), future).await {
Ok(result) => result,
Err(_) => {
connections.record_timeout(lane);
return Err(crate::DataLayerError::TimedOut(format!(
"{operation} exceeded {timeout_ms}ms timeout"
)));
}
}
} else {
future.await
};
if result.is_err() {
connections.record_error(lane);
}
result
}

View File

@@ -0,0 +1,699 @@
use std::time::Duration;
use crate::error::RedisResultExt;
use crate::redis::{
cmd, run_lane_with_timeout, script, RedisCmd, RedisConnectionLane, RedisConnectionRouter,
RedisKeyspace, RedisLaneDiagnostics,
};
use crate::{
DataLayerError, RateLimitCheck, RateLimitInput, RateLimitScope, RuntimeSemaphoreError,
};
const RATE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#"
local user_key = KEYS[1]
local key_key = KEYS[2]
local user_limit = tonumber(ARGV[1])
local key_limit = tonumber(ARGV[2])
local ttl = tonumber(ARGV[3])
local user_count = 0
if user_limit > 0 then
user_count = tonumber(redis.call('GET', user_key) or '0')
if user_count >= user_limit then
return {0, 1, user_limit, 0}
end
end
local key_count = 0
if key_limit > 0 then
key_count = tonumber(redis.call('GET', key_key) or '0')
if key_count >= key_limit then
return {0, 2, key_limit, 0}
end
end
local remaining = -1
if user_limit > 0 then
user_count = redis.call('INCR', user_key)
redis.call('EXPIRE', user_key, ttl)
remaining = user_limit - user_count
end
if key_limit > 0 then
key_count = redis.call('INCR', key_key)
redis.call('EXPIRE', key_key, ttl)
local key_remaining = key_limit - key_count
if remaining == -1 or key_remaining < remaining then
remaining = key_remaining
end
end
return {1, 0, 0, remaining}
"#;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub struct RedisRuntimeDiagnostics {
pub connected_clients: Option<u64>,
pub total_connections_received: Option<u64>,
pub lanes: Vec<RedisLaneDiagnostics>,
}
#[derive(Debug, Clone)]
pub(crate) struct RedisRuntimeRunner {
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
command_timeout_ms: Option<u64>,
}
impl RedisRuntimeRunner {
pub(crate) fn new(
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
command_timeout_ms: Option<u64>,
) -> Self {
Self {
connections,
keyspace,
command_timeout_ms,
}
}
pub(crate) async fn ping(&self) -> Result<(), DataLayerError> {
let pong = self
.query_string(RedisConnectionLane::Fast, "runtime redis ping", cmd("PING"))
.await?;
if pong.eq_ignore_ascii_case("PONG") {
Ok(())
} else {
Err(DataLayerError::UnexpectedValue(format!(
"unexpected runtime redis ping response {pong}"
)))
}
}
pub(crate) async fn diagnostics(&self) -> Result<RedisRuntimeDiagnostics, DataLayerError> {
let info = self
.query_string(
RedisConnectionLane::Admin,
"runtime redis diagnostics",
cmd("INFO"),
)
.await?;
Ok(parse_diagnostics(
&info,
self.connections.lane_diagnostics(),
))
}
pub(crate) async fn kv_set_plain(
&self,
key: &str,
value: String,
) -> Result<(), DataLayerError> {
let namespaced_key = self.keyspace.key(key);
let mut command = cmd("SET");
command.arg(namespaced_key).arg(value);
self.query_string(RedisConnectionLane::Fast, "runtime kv set", command)
.await?;
Ok(())
}
pub(crate) async fn kv_set_with_ttl(
&self,
key: &str,
value: String,
ttl: Duration,
) -> Result<(), DataLayerError> {
let namespaced_key = self.keyspace.key(key);
let mut command = cmd("PSETEX");
command
.arg(namespaced_key)
.arg(u64::try_from(ttl.as_millis().max(1)).unwrap_or(u64::MAX))
.arg(value);
self.query_string(RedisConnectionLane::Fast, "runtime kv set ttl", command)
.await?;
Ok(())
}
pub(crate) async fn kv_get_many(
&self,
keys: &[String],
) -> Result<Vec<Option<String>>, DataLayerError> {
let namespaced = keys
.iter()
.map(|key| self.keyspace.key(key))
.collect::<Vec<_>>();
let mut command = cmd("MGET");
command.arg(&namespaced);
self.query(RedisConnectionLane::Fast, "runtime kv mget", command)
.await
}
pub(crate) async fn kv_delete_many(&self, keys: &[String]) -> Result<usize, DataLayerError> {
let prefix = self.keyspace.key("");
let namespaced = keys
.iter()
.map(|key| {
if key_belongs_to_prefix(key, &prefix) {
key.clone()
} else {
self.keyspace.key(key)
}
})
.collect::<Vec<_>>();
let mut command = cmd("DEL");
command.arg(&namespaced);
let deleted = self
.query_i64(
RedisConnectionLane::Admin,
"runtime kv delete many",
command,
)
.await?;
Ok(usize::try_from(deleted).unwrap_or(0))
}
pub(crate) async fn kv_ttl_seconds(&self, key: &str) -> Result<Option<i64>, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
let mut command = cmd("TTL");
command.arg(&namespaced_key);
let ttl = self
.query_i64(RedisConnectionLane::Fast, "runtime kv ttl", command)
.await?;
Ok((ttl >= -1).then_some(ttl))
}
pub(crate) async fn scan_keys(
&self,
pattern: &str,
count: usize,
) -> Result<Vec<String>, DataLayerError> {
let pattern = self.keyspace.key(pattern);
run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Admin,
self.command_timeout_ms,
"runtime scan keys",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Admin);
let mut cursor = 0u64;
let mut keys = Vec::new();
loop {
let (next_cursor, mut batch) = cmd("SCAN")
.arg(cursor)
.arg("MATCH")
.arg(&pattern)
.arg("COUNT")
.arg(count.max(1))
.query_async::<(u64, Vec<String>)>(&mut connection)
.await
.map_redis_err()?;
keys.append(&mut batch);
if next_cursor == 0 {
break;
}
cursor = next_cursor;
}
keys.sort();
Ok(keys)
},
)
.await
}
pub(crate) async fn check_and_consume_rate_limit(
&self,
input: RateLimitInput<'_>,
) -> Result<RateLimitCheck, DataLayerError> {
let user_key = self.keyspace.key(input.user_key);
let key_key = self.keyspace.key(input.key_key);
let raw = run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Fast,
self.command_timeout_ms,
"runtime rate limit check",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
script(RATE_LIMIT_CHECK_AND_CONSUME_SCRIPT)
.key(user_key)
.key(key_key)
.arg(i64::from(input.user_limit))
.arg(i64::from(input.key_limit))
.arg(i64::try_from(input.ttl_seconds.max(1)).unwrap_or(i64::MAX))
.invoke_async::<Vec<i64>>(&mut connection)
.await
.map_redis_err()
},
)
.await?;
if raw.first().copied().unwrap_or_default() == 1 {
return Ok(RateLimitCheck::Allowed {
remaining: raw
.get(3)
.copied()
.and_then(|value| u32::try_from(value).ok())
.unwrap_or_default(),
});
}
let scope = match raw.get(1).copied().unwrap_or_default() {
2 => RateLimitScope::Key,
_ => RateLimitScope::User,
};
let limit = raw
.get(2)
.copied()
.and_then(|value| u32::try_from(value).ok())
.unwrap_or(match scope {
RateLimitScope::User => input.user_limit,
RateLimitScope::Key => input.key_limit,
});
Ok(RateLimitCheck::Rejected { scope, limit })
}
pub(crate) async fn set_add(&self, key: &str, member: &str) -> Result<bool, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("SADD");
command.arg(&key).arg(member);
Ok(self
.query_i64(RedisConnectionLane::Fast, "runtime set add", command)
.await?
> 0)
}
pub(crate) async fn set_remove(&self, key: &str, member: &str) -> Result<bool, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("SREM");
command.arg(&key).arg(member);
Ok(self
.query_i64(RedisConnectionLane::Fast, "runtime set remove", command)
.await?
> 0)
}
pub(crate) async fn set_members(&self, key: &str) -> Result<Vec<String>, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("SMEMBERS");
command.arg(&key);
let mut values = self
.query::<Vec<String>>(RedisConnectionLane::Admin, "runtime set members", command)
.await?;
values.sort();
Ok(values)
}
pub(crate) async fn set_len(&self, key: &str) -> Result<usize, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("SCARD");
command.arg(&key);
let len = self
.query_i64(RedisConnectionLane::Fast, "runtime set len", command)
.await?;
Ok(usize::try_from(len).unwrap_or(0))
}
pub(crate) async fn score_set(
&self,
key: &str,
member: &str,
score: f64,
) -> Result<(), DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("ZADD");
command.arg(&key).arg(score).arg(member);
self.query_i64(RedisConnectionLane::Fast, "runtime score set", command)
.await?;
Ok(())
}
pub(crate) async fn score_many(
&self,
key: &str,
members: &[String],
) -> Result<Vec<Option<f64>>, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("ZMSCORE");
command.arg(&key);
for member in members {
command.arg(member);
}
self.query(RedisConnectionLane::Fast, "runtime score many", command)
.await
}
pub(crate) async fn score_range_by_min(
&self,
key: &str,
min_score: f64,
) -> Result<Vec<String>, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("ZRANGEBYSCORE");
command.arg(&key).arg(min_score).arg("+inf");
self.query(RedisConnectionLane::Admin, "runtime score range", command)
.await
}
pub(crate) async fn score_remove_by_score(
&self,
key: &str,
max_score: f64,
) -> Result<usize, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("ZREMRANGEBYSCORE");
command.arg(&key).arg("-inf").arg(max_score);
let removed = self
.query_i64(RedisConnectionLane::Admin, "runtime score trim", command)
.await?;
Ok(usize::try_from(removed).unwrap_or(0))
}
pub(crate) async fn score_remove(
&self,
key: &str,
member: &str,
) -> Result<bool, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("ZREM");
command.arg(&key).arg(member);
Ok(self
.query_i64(RedisConnectionLane::Fast, "runtime score remove", command)
.await?
> 0)
}
pub(crate) async fn score_remove_by_rank(
&self,
key: &str,
start: i64,
stop: i64,
) -> Result<usize, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("ZREMRANGEBYRANK");
command.arg(&key).arg(start).arg(stop);
let removed = self
.query_i64(
RedisConnectionLane::Admin,
"runtime score rank trim",
command,
)
.await?;
Ok(usize::try_from(removed).unwrap_or(0))
}
pub(crate) async fn score_len(&self, key: &str) -> Result<usize, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("ZCARD");
command.arg(&key);
let len = self
.query_i64(RedisConnectionLane::Fast, "runtime score len", command)
.await?;
Ok(usize::try_from(len).unwrap_or(0))
}
pub(crate) async fn key_expire(
&self,
key: &str,
ttl: Duration,
) -> Result<bool, DataLayerError> {
let key = self.keyspace.key(key);
let mut command = cmd("PEXPIRE");
command
.arg(&key)
.arg(u64::try_from(ttl.as_millis()).unwrap_or(u64::MAX));
Ok(self
.query_i64(RedisConnectionLane::Fast, "runtime key expire", command)
.await?
> 0)
}
pub(crate) async fn semaphore_try_acquire(
&self,
gate: &'static str,
limit: usize,
key: &str,
token: &str,
lease_ttl_ms: u64,
timeout_ms: Option<u64>,
) -> Result<(i64, i64), RuntimeSemaphoreError> {
let now_ms = crate::unix_time_ms();
let expires_at_ms = now_ms.saturating_add(lease_ttl_ms);
let key = self.keyspace.key(key);
let timeout_ms = timeout_ms.or(self.command_timeout_ms);
run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Fast,
timeout_ms,
"runtime semaphore acquire",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
script(
"redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]); \
local count = redis.call('ZCARD', KEYS[1]); \
if count >= tonumber(ARGV[3]) then \
redis.call('PEXPIRE', KEYS[1], ARGV[5]); \
return {0, count}; \
end; \
redis.call('ZADD', KEYS[1], ARGV[2], ARGV[4]); \
count = redis.call('ZCARD', KEYS[1]); \
redis.call('PEXPIRE', KEYS[1], ARGV[5]); \
return {1, count};",
)
.key(&key)
.arg(now_ms as i64)
.arg(expires_at_ms as i64)
.arg(limit as i64)
.arg(token)
.arg(lease_ttl_ms as i64)
.invoke_async::<(i64, i64)>(&mut connection)
.await
.map_redis_err()
},
)
.await
.map_err(|err| RuntimeSemaphoreError::Unavailable {
gate,
limit,
message: format!("acquire failed: {err}"),
})
}
pub(crate) async fn semaphore_renew(
&self,
gate: &'static str,
limit: usize,
key: &str,
token: &str,
lease_ttl_ms: u64,
timeout_ms: Option<u64>,
) -> Result<i64, RuntimeSemaphoreError> {
let now_ms = crate::unix_time_ms();
let expires_at_ms = now_ms.saturating_add(lease_ttl_ms);
let key = self.keyspace.key(key);
let timeout_ms = timeout_ms.or(self.command_timeout_ms);
run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Fast,
timeout_ms,
"runtime semaphore renew",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
script(
"redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]); \
local score = redis.call('ZSCORE', KEYS[1], ARGV[2]); \
if not score then return 0; end; \
redis.call('ZADD', KEYS[1], 'XX', ARGV[3], ARGV[2]); \
redis.call('PEXPIRE', KEYS[1], ARGV[4]); \
return 1;",
)
.key(&key)
.arg(now_ms as i64)
.arg(token)
.arg(expires_at_ms as i64)
.arg(lease_ttl_ms as i64)
.invoke_async::<i64>(&mut connection)
.await
.map_redis_err()
},
)
.await
.map_err(|err| RuntimeSemaphoreError::Unavailable {
gate,
limit,
message: format!("renew failed: {err}"),
})
}
pub(crate) async fn semaphore_release(
&self,
gate: &'static str,
limit: usize,
key: &str,
token: &str,
timeout_ms: Option<u64>,
) -> Result<(), RuntimeSemaphoreError> {
let key = self.keyspace.key(key);
let timeout_ms = timeout_ms.or(self.command_timeout_ms);
run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Fast,
timeout_ms,
"runtime semaphore release",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
script(
"local removed = redis.call('ZREM', KEYS[1], ARGV[1]); \
if removed > 0 and redis.call('ZCARD', KEYS[1]) == 0 then \
redis.call('DEL', KEYS[1]); \
end; \
return removed;",
)
.key(&key)
.arg(token)
.invoke_async::<i64>(&mut connection)
.await
.map(|_| ())
.map_redis_err()
},
)
.await
.map_err(|err| RuntimeSemaphoreError::Unavailable {
gate,
limit,
message: format!("release failed: {err}"),
})
}
pub(crate) async fn semaphore_live_count(
&self,
gate: &'static str,
limit: usize,
key: &str,
timeout_ms: Option<u64>,
) -> Result<usize, RuntimeSemaphoreError> {
let now_ms = crate::unix_time_ms();
let key = self.keyspace.key(key);
let timeout_ms = timeout_ms.or(self.command_timeout_ms);
run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Fast,
timeout_ms,
"runtime semaphore snapshot",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
script(
"redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]); \
return redis.call('ZCARD', KEYS[1]);",
)
.key(&key)
.arg(now_ms as i64)
.invoke_async::<i64>(&mut connection)
.await
.map(|value| value.max(0) as usize)
.map_redis_err()
},
)
.await
.map_err(|err| RuntimeSemaphoreError::Unavailable {
gate,
limit,
message: format!("snapshot failed: {err}"),
})
}
async fn query<T>(
&self,
lane: RedisConnectionLane,
operation: &'static str,
command: RedisCmd,
) -> Result<T, DataLayerError>
where
T: redis::FromRedisValue,
{
run_lane_with_timeout(
&self.connections,
lane,
self.command_timeout_ms,
operation,
async {
let mut connection = self.connections.connection(lane);
command
.query_async::<T>(&mut connection)
.await
.map_redis_err()
},
)
.await
}
async fn query_i64(
&self,
lane: RedisConnectionLane,
operation: &'static str,
command: RedisCmd,
) -> Result<i64, DataLayerError> {
self.query(lane, operation, command).await
}
async fn query_string(
&self,
lane: RedisConnectionLane,
operation: &'static str,
command: RedisCmd,
) -> Result<String, DataLayerError> {
self.query(lane, operation, command).await
}
}
fn parse_diagnostics(info: &str, lanes: Vec<RedisLaneDiagnostics>) -> RedisRuntimeDiagnostics {
RedisRuntimeDiagnostics {
connected_clients: parse_info_u64(info, "connected_clients"),
total_connections_received: parse_info_u64(info, "total_connections_received"),
lanes,
}
}
fn parse_info_u64(info: &str, key: &str) -> Option<u64> {
info.lines().find_map(|line| {
let (name, value) = line.split_once(':')?;
(name == key)
.then(|| value.trim().parse::<u64>().ok())
.flatten()
})
}
fn key_belongs_to_prefix(key: &str, prefix: &str) -> bool {
prefix.is_empty()
|| key == prefix
|| key
.strip_prefix(prefix)
.is_some_and(|rest| rest.starts_with(':'))
}
#[cfg(test)]
mod tests {
use super::{key_belongs_to_prefix, parse_diagnostics, RedisRuntimeDiagnostics};
#[test]
fn parses_runtime_diagnostics_from_info() {
let parsed = parse_diagnostics(
"# Clients\r\nconnected_clients:5\r\n# Stats\r\ntotal_connections_received:42\r\n",
Vec::new(),
);
assert_eq!(
parsed,
RedisRuntimeDiagnostics {
connected_clients: Some(5),
total_connections_received: Some(42),
lanes: Vec::new(),
}
);
}
#[test]
fn detects_namespaced_key_prefix_on_boundary() {
assert!(key_belongs_to_prefix("aether:cache:item", "aether"));
assert!(key_belongs_to_prefix("aether", "aether"));
assert!(key_belongs_to_prefix("raw:key", ""));
assert!(!key_belongs_to_prefix("aetherish:cache:item", "aether"));
}
}

View File

@@ -1,13 +1,15 @@
use std::collections::BTreeMap;
use std::future::Future;
use std::time::Duration;
use redis::from_redis_value;
use redis::streams::StreamReadReply;
use redis::Value as RedisValue;
use crate::error::{redis_error, RedisResultExt};
use crate::redis::{RedisClient, RedisKeyspace};
use crate::redis::{
run_lane_with_timeout, RedisClientConfig, RedisClientFactory, RedisConnectionLane,
RedisConnectionRouter, RedisKeyspace,
};
use crate::DataLayerError;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
@@ -113,27 +115,35 @@ impl RedisStreamRunnerConfig {
#[derive(Debug, Clone)]
pub struct RedisStreamRunner {
client: RedisClient,
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
config: RedisStreamRunnerConfig,
}
impl RedisStreamRunner {
pub fn new(
client: RedisClient,
pub(crate) fn new(
connections: RedisConnectionRouter,
keyspace: RedisKeyspace,
config: RedisStreamRunnerConfig,
) -> Result<Self, DataLayerError> {
config.validate()?;
Ok(Self {
client,
connections,
keyspace,
config,
})
}
pub fn client(&self) -> &RedisClient {
&self.client
pub async fn from_config(
config: RedisClientConfig,
runner_config: RedisStreamRunnerConfig,
) -> Result<Self, DataLayerError> {
let factory = RedisClientFactory::new(config)?;
let keyspace = factory.config().keyspace();
let connections = factory
.connect_router(runner_config.command_timeout_ms)
.await?;
Self::new(connections, keyspace, runner_config)
}
pub fn keyspace(&self) -> &RedisKeyspace {
@@ -144,6 +154,13 @@ impl RedisStreamRunner {
self.config
}
pub(crate) fn with_config(
&self,
config: RedisStreamRunnerConfig,
) -> Result<Self, DataLayerError> {
Self::new(self.connections.clone(), self.keyspace.clone(), config)
}
pub async fn ensure_consumer_group(
&self,
stream: &RedisStreamName,
@@ -154,27 +171,27 @@ impl RedisStreamRunner {
validate_group(group)?;
validate_stream_position(start_id)?;
self.run_with_timeout("redis stream ensure consumer group", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
let result = redis::cmd("XGROUP")
.arg("CREATE")
.arg(&stream.0)
.arg(&group.0)
.arg(start_id)
.arg("MKSTREAM")
.query_async::<String>(&mut connection)
.await;
self.run_with_timeout(
RedisConnectionLane::Stream,
"redis stream ensure consumer group",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Stream);
let result = redis::cmd("XGROUP")
.arg("CREATE")
.arg(&stream.0)
.arg(&group.0)
.arg(start_id)
.arg("MKSTREAM")
.query_async::<String>(&mut connection)
.await;
match result {
Ok(_) => Ok(()),
Err(err) if err.code() == Some("BUSYGROUP") => Ok(()),
Err(err) => Err(redis_error(err)),
}
})
match result {
Ok(_) => Ok(()),
Err(err) if err.code() == Some("BUSYGROUP") => Ok(()),
Err(err) => Err(redis_error(err)),
}
},
)
.await
}
@@ -199,12 +216,8 @@ impl RedisStreamRunner {
));
}
self.run_with_timeout("redis stream append", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Stream, "redis stream append", async {
let mut connection = self.connections.connection(RedisConnectionLane::Stream);
let mut command = redis::cmd("XADD");
command.arg(&stream.0);
if let Some(maxlen) = maxlen.filter(|value| *value > 0) {
@@ -256,12 +269,13 @@ impl RedisStreamRunner {
validate_group(group)?;
validate_consumer(consumer)?;
self.run_with_timeout("redis stream read group", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
let lane = if self.config.read_block_ms.is_some() {
RedisConnectionLane::BlockingStream
} else {
RedisConnectionLane::Stream
};
self.run_with_timeout(lane, "redis stream read group", async {
let mut connection = self.connections.connection(lane);
let mut command = redis::cmd("XREADGROUP");
command
.arg("GROUP")
@@ -312,12 +326,8 @@ impl RedisStreamRunner {
return Ok(0);
}
self.run_with_timeout("redis stream ack", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Stream, "redis stream ack", async {
let mut connection = self.connections.connection(RedisConnectionLane::Stream);
let mut command = redis::cmd("XACK");
command.arg(&stream.0).arg(&group.0);
for id in ids {
@@ -341,12 +351,8 @@ impl RedisStreamRunner {
return Ok(0);
}
self.run_with_timeout("redis stream delete", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Stream, "redis stream delete", async {
let mut connection = self.connections.connection(RedisConnectionLane::Stream);
let mut command = redis::cmd("XDEL");
command.arg(&stream.0);
for id in ids {
@@ -374,12 +380,8 @@ impl RedisStreamRunner {
validate_stream_position(start_id)?;
config.validate()?;
self.run_with_timeout("redis stream reclaim", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
self.run_with_timeout(RedisConnectionLane::Stream, "redis stream reclaim", async {
let mut connection = self.connections.connection(RedisConnectionLane::Stream);
let reply = redis::cmd("XAUTOCLAIM")
.arg(&stream.0)
.arg(&group.0)
@@ -399,21 +401,21 @@ impl RedisStreamRunner {
async fn run_with_timeout<T, F>(
&self,
lane: RedisConnectionLane,
operation: &'static str,
future: F,
) -> Result<T, DataLayerError>
where
F: Future<Output = Result<T, DataLayerError>>,
{
if let Some(timeout_ms) = self.config.command_timeout_ms {
tokio::time::timeout(Duration::from_millis(timeout_ms), future)
.await
.map_err(|_| {
DataLayerError::TimedOut(format!("{operation} exceeded {timeout_ms}ms timeout"))
})?
} else {
future.await
}
run_lane_with_timeout(
&self.connections,
lane,
self.config.command_timeout_ms,
operation,
future,
)
.await
}
}
@@ -571,31 +573,12 @@ mod tests {
use std::collections::BTreeMap;
use super::{
parse_reclaim_result, RedisConsumerGroup, RedisConsumerName, RedisStreamName,
RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunner,
RedisStreamRunnerConfig,
parse_reclaim_result, validate_consumer, validate_group, validate_stream_name,
validate_stream_position, RedisConsumerName, RedisStreamName, RedisStreamReclaimConfig,
RedisStreamReclaimResult, RedisStreamRunnerConfig,
};
use crate::redis::{RedisClientConfig, RedisClientFactory};
use redis::Value as RedisValue;
fn sample_runner() -> RedisStreamRunner {
let config = RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
};
let client = RedisClientFactory::new(config.clone())
.expect("factory should build")
.connect_lazy()
.expect("client should build");
RedisStreamRunner::new(
client,
config.keyspace(),
RedisStreamRunnerConfig::default(),
)
.expect("runner should build")
}
#[test]
fn validates_stream_runner_config() {
assert!(RedisStreamRunnerConfig {
@@ -643,55 +626,17 @@ mod tests {
#[test]
fn runner_reuses_client_and_keyspace() {
let runner = sample_runner();
assert_eq!(runner.config(), RedisStreamRunnerConfig::default());
assert_eq!(
runner.keyspace().stream_name("audit").0,
"aether:stream:audit"
);
let _client_ref = runner.client();
RedisStreamRunnerConfig::default()
.validate()
.expect("default stream config should be valid");
}
#[tokio::test]
async fn rejects_invalid_inputs_before_network() {
let runner = sample_runner();
let stream = RedisStreamName("aether:stream:audit".to_string());
let group = RedisConsumerGroup("audit-workers".to_string());
let consumer = RedisConsumerName("worker-1".to_string());
assert!(runner
.ensure_consumer_group(&stream, &group, "")
.await
.is_err());
assert!(runner
.append_fields(&stream, &BTreeMap::new())
.await
.is_err());
assert!(runner
.append_json(&stream, "", &serde_json::json!({"ok": true}))
.await
.is_err());
assert!(runner
.read_group(&stream, &group, &RedisConsumerName(String::new()))
.await
.is_err());
assert_eq!(
runner.ack(&stream, &group, &[]).await.expect("empty ack"),
0
);
assert_eq!(runner.delete(&stream, &[]).await.expect("empty delete"), 0);
assert!(runner
.claim_stale(
&stream,
&group,
&consumer,
"",
RedisStreamReclaimConfig::default()
)
.await
.is_err());
let _ = consumer;
#[test]
fn rejects_invalid_inputs_before_network() {
assert!(validate_stream_name(&RedisStreamName(String::new())).is_err());
assert!(validate_group(&super::RedisConsumerGroup(String::new())).is_err());
assert!(validate_consumer(&RedisConsumerName(String::new())).is_err());
assert!(validate_stream_position("").is_err());
}
#[test]