use super::{
JobIndexPage, JobIndexRow, JobIndexWaitStats, RangeItem, RangeOrder, RangePage, Storage,
StorageOperation,
};
use crate::UtcDateTime;
use std::collections::{BTreeMap, HashMap};
use std::ops::Bound::{Excluded, Unbounded};
use std::time::{Duration, Instant};
const METRIC_BUCKET_MS: i64 = 60_000;
fn minute_bucket(at: UtcDateTime) -> i64 {
at.timestamp_millis().div_euclid(METRIC_BUCKET_MS)
}
#[derive(Clone, Default)]
struct StoredValue {
value: Vec<u8>,
expires_at: Option<Instant>,
}
#[derive(Clone, Default)]
struct MemoryRange {
next_cursor: i64,
values: BTreeMap<i64, RangeValue>,
cursors: HashMap<Vec<u8>, i64>,
}
#[derive(Clone)]
struct RangeValue {
value: Vec<u8>,
expires_at: Option<Instant>,
}
#[derive(Clone, Default)]
struct MemoryState {
values: HashMap<String, StoredValue>,
ranges: HashMap<String, MemoryRange>,
job_index: HashMap<(String, String), JobIndexRow>,
transition_metrics: HashMap<(String, String, i64), usize>,
wait_metrics: HashMap<(String, String, i64), (usize, i64, i64)>,
}
#[derive(Default)]
pub struct MemoryStorage {
state: async_lock::Mutex<MemoryState>,
}
impl MemoryStorage {
pub fn new() -> Self {
Self::default()
}
}
#[async_trait::async_trait]
impl Storage for MemoryStorage {
async fn get(&self, key: &str) -> anyhow::Result<Option<Vec<u8>>> {
let mut state = self.state.lock().await;
let expired = state
.values
.get(key)
.and_then(|value| value.expires_at)
.is_some_and(|expires_at| expires_at <= Instant::now());
if expired {
state.values.remove(key);
return Ok(None);
}
Ok(state.values.get(key).map(|value| value.value.clone()))
}
async fn set_if_absent(&self, key: &str, value: &[u8]) -> anyhow::Result<bool> {
let mut state = self.state.lock().await;
let expired = state
.values
.get(key)
.and_then(|value| value.expires_at)
.is_some_and(|expires_at| expires_at <= Instant::now());
if !expired && state.values.contains_key(key) {
return Ok(false);
}
state.values.insert(
key.to_string(),
StoredValue {
value: value.to_vec(),
expires_at: None,
},
);
Ok(true)
}
async fn apply(&self, operations: Vec<StorageOperation>) -> anyhow::Result<()> {
let mut state = self.state.lock().await;
let mut next_state = state.clone();
for operation in operations {
match operation {
StorageOperation::Set { key, value } => {
next_state.values.insert(
key,
StoredValue {
value,
expires_at: None,
},
);
}
StorageOperation::Delete { key } => {
next_state.values.remove(&key);
}
StorageOperation::Expire { key, ttl_seconds } => {
if let Some(value) = next_state.values.get_mut(&key) {
value.expires_at = Some(
Instant::now()
.checked_add(Duration::from_secs(u64::try_from(ttl_seconds)?))
.ok_or_else(|| anyhow::anyhow!("memory expiry is out of range"))?,
);
}
}
StorageOperation::RangeAdd { key, value } => {
let range = next_state.ranges.entry(key).or_default();
if let Some(cursor) = range.cursors.get(&value) {
if let Some(existing) = range.values.get_mut(cursor) {
existing.expires_at = None;
}
} else {
range.next_cursor = range
.next_cursor
.checked_add(1)
.ok_or_else(|| anyhow::anyhow!("memory range cursor overflow"))?;
range.values.insert(
range.next_cursor,
RangeValue {
value: value.clone(),
expires_at: None,
},
);
range.cursors.insert(value, range.next_cursor);
}
}
StorageOperation::RangeRemove { key, value } => {
if let Some(range) = next_state.ranges.get_mut(&key) {
if let Some(cursor) = range.cursors.remove(&value) {
range.values.remove(&cursor);
}
}
}
StorageOperation::RangeExpire {
key,
value,
ttl_seconds,
} => {
if let Some(range) = next_state.ranges.get_mut(&key) {
if let Some(cursor) = range.cursors.get(&value) {
if let Some(existing) = range.values.get_mut(cursor) {
existing.expires_at = Some(
Instant::now()
.checked_add(Duration::from_secs(u64::try_from(
ttl_seconds,
)?))
.ok_or_else(|| {
anyhow::anyhow!("memory range expiry is out of range")
})?,
);
}
}
}
}
StorageOperation::RangeClear { key } => {
next_state.ranges.remove(&key);
}
}
}
*state = next_state;
Ok(())
}
async fn range_page(
&self,
key: &str,
cursor: Option<i64>,
limit: usize,
order: RangeOrder,
) -> anyhow::Result<RangePage> {
let state = self.state.lock().await;
let Some(range) = state.ranges.get(key) else {
return Ok(RangePage {
items: Vec::new(),
next_cursor: None,
});
};
let mut items = Vec::with_capacity(limit);
match order {
RangeOrder::OldestFirst => {
let start = cursor.unwrap_or(i64::MIN);
for (&item_cursor, value) in range
.values
.range((Excluded(start), Unbounded))
.filter(|(_, value)| {
value
.expires_at
.is_none_or(|expires_at| expires_at > Instant::now())
})
.take(limit)
{
items.push(RangeItem {
cursor: item_cursor,
value: value.value.clone(),
});
}
}
RangeOrder::NewestFirst => {
if let Some(end) = cursor {
for (&item_cursor, value) in range
.values
.range((Unbounded, Excluded(end)))
.rev()
.filter(|(_, value)| {
value
.expires_at
.is_none_or(|expires_at| expires_at > Instant::now())
})
.take(limit)
{
items.push(RangeItem {
cursor: item_cursor,
value: value.value.clone(),
});
}
} else {
for (&item_cursor, value) in range
.values
.iter()
.rev()
.filter(|(_, value)| {
value
.expires_at
.is_none_or(|expires_at| expires_at > Instant::now())
})
.take(limit)
{
items.push(RangeItem {
cursor: item_cursor,
value: value.value.clone(),
});
}
}
}
}
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> {
let mut state = self.state.lock().await;
let now = Instant::now();
let count = if let Some(range) = state.ranges.get_mut(key) {
range
.values
.retain(|_, value| value.expires_at.is_none_or(|expires_at| expires_at > now));
range.values.len()
} else {
0
};
Ok(count)
}
async fn job_index_upsert(&self, namespace: &str, mut row: JobIndexRow) -> anyhow::Result<()> {
let mut state = self.state.lock().await;
let key = (namespace.to_string(), row.job_id.clone());
if let Some(existing) = state.job_index.get(&key) {
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();
}
}
state.job_index.insert(key, row);
Ok(())
}
async fn job_index_stage_counts(
&self,
namespace: &str,
) -> anyhow::Result<std::collections::HashMap<String, usize>> {
let state = self.state.lock().await;
let now = chrono::Utc::now();
let mut counts = std::collections::HashMap::new();
for row in state
.job_index
.iter()
.filter(|((ns, _), row)| ns == namespace && !is_expired(row, now))
.map(|(_, row)| row)
{
*counts.entry(row.stage.clone()).or_insert(0) += 1;
}
Ok(counts)
}
async fn job_index_list_by_stage(
&self,
namespace: &str,
stage: &str,
cursor: Option<i64>,
limit: usize,
) -> anyhow::Result<JobIndexPage> {
let state = self.state.lock().await;
let now = chrono::Utc::now();
let mut rows: Vec<&JobIndexRow> = state
.job_index
.iter()
.filter(|((ns, _), row)| ns == namespace && row.stage == stage && !is_expired(row, now))
.map(|(_, row)| row)
.filter(|row| cursor.is_none_or(|cursor| row.stage_date.timestamp_millis() < cursor))
.collect();
rows.sort_by(|a, b| b.stage_date.cmp(&a.stage_date));
rows.truncate(limit);
let next_cursor = (rows.len() == limit)
.then(|| rows.last().map(|row| row.stage_date.timestamp_millis()))
.flatten();
Ok(JobIndexPage {
items: rows.into_iter().cloned().collect(),
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 state = self.state.lock().await;
let partition = i64::from(partition);
let mut rows: Vec<&JobIndexRow> = state
.job_index
.iter()
.filter(|((ns, _), row)| {
ns == namespace
&& row.topic.as_deref() == Some(topic)
&& row.partition == Some(partition)
})
.map(|(_, row)| row)
.filter(|row| cursor.is_none_or(|cursor| row.sequence.is_some_and(|s| s < cursor)))
.collect();
rows.sort_by(|a, b| b.sequence.cmp(&a.sequence));
rows.truncate(limit);
let next_cursor = (rows.len() == limit)
.then(|| rows.last().and_then(|row| row.sequence))
.flatten();
Ok(JobIndexPage {
items: rows.into_iter().cloned().collect(),
next_cursor,
})
}
async fn job_index_partition_neighbor(
&self,
namespace: &str,
topic: &str,
partition: u32,
sequence: i64,
older: bool,
) -> anyhow::Result<Option<JobIndexRow>> {
let state = self.state.lock().await;
let partition = i64::from(partition);
let target = if older { sequence - 1 } else { sequence + 1 };
Ok(state
.job_index
.iter()
.find(|((ns, _), row)| {
ns == namespace
&& row.topic.as_deref() == Some(topic)
&& row.partition == Some(partition)
&& row.sequence == Some(target)
})
.map(|(_, row)| row.clone()))
}
async fn job_index_list_continuations(
&self,
namespace: &str,
parent_job_id: &str,
limit: usize,
) -> anyhow::Result<Vec<JobIndexRow>> {
let state = self.state.lock().await;
let mut rows: Vec<&JobIndexRow> = state
.job_index
.iter()
.filter(|((ns, _), row)| {
ns == namespace && row.parent_job_id.as_deref() == Some(parent_job_id)
})
.map(|(_, row)| row)
.collect();
rows.sort_by(|a, b| a.stage_date.cmp(&b.stage_date));
rows.truncate(limit);
Ok(rows.into_iter().cloned().collect())
}
async fn job_index_truncate_partition(
&self,
namespace: &str,
topic: &str,
partition: u32,
keep_last: usize,
) -> anyhow::Result<usize> {
let mut state = self.state.lock().await;
let partition_id = i64::from(partition);
let mut keys: Vec<(String, String)> = state
.job_index
.iter()
.filter(|((ns, _), row)| {
ns == namespace
&& row.topic.as_deref() == Some(topic)
&& row.partition == Some(partition_id)
})
.map(|(key, _)| key.clone())
.collect();
keys.sort_by_key(|key| std::cmp::Reverse(state.job_index[key].sequence));
let to_remove = keys.split_off(keep_last.min(keys.len()));
let removed = to_remove.len();
for key in to_remove {
state.job_index.remove(&key);
}
Ok(removed)
}
async fn job_index_record_transition(
&self,
namespace: &str,
stage: &str,
at: UtcDateTime,
) -> anyhow::Result<()> {
let mut state = self.state.lock().await;
*state
.transition_metrics
.entry((namespace.to_string(), stage.to_string(), minute_bucket(at)))
.or_insert(0) += 1;
Ok(())
}
async fn job_index_record_metrics_batch(
&self,
namespace: &str,
batch: &crate::storage::JobIndexMetricsBatch,
) -> anyhow::Result<()> {
let mut state = self.state.lock().await;
for (stage, at, count) in &batch.transitions {
*state
.transition_metrics
.entry((namespace.to_string(), stage.clone(), minute_bucket(*at)))
.or_insert(0) += usize::try_from(*count)?;
}
for (mode, at, count, sum_ms, max_ms) in &batch.waits {
let entry = state
.wait_metrics
.entry((namespace.to_string(), mode.clone(), minute_bucket(*at)))
.or_insert((0, 0, 0));
entry.0 += usize::try_from(*count)?;
entry.1 += sum_ms;
entry.2 = entry.2.max(*max_ms);
}
Ok(())
}
async fn job_index_recent_transition_count(
&self,
namespace: &str,
stage: &str,
now: UtcDateTime,
) -> anyhow::Result<usize> {
let state = self.state.lock().await;
let bucket = minute_bucket(now);
Ok(state
.transition_metrics
.iter()
.filter(|((ns, s, b), _)| ns == namespace && s == stage && *b >= bucket - 1)
.map(|(_, count)| *count)
.sum())
}
async fn job_index_record_wait(
&self,
namespace: &str,
mode: &str,
wait_ms: i64,
at: UtcDateTime,
) -> anyhow::Result<()> {
let mut state = self.state.lock().await;
let entry = state
.wait_metrics
.entry((namespace.to_string(), mode.to_string(), minute_bucket(at)))
.or_insert((0, 0, 0));
entry.0 += 1;
entry.1 += wait_ms;
entry.2 = entry.2.max(wait_ms);
Ok(())
}
async fn job_index_recent_wait_stats(
&self,
namespace: &str,
mode: &str,
now: UtcDateTime,
) -> anyhow::Result<JobIndexWaitStats> {
let state = self.state.lock().await;
let bucket = minute_bucket(now);
let (count, sum_ms, max_ms) = state
.wait_metrics
.iter()
.filter(|((ns, m, b), _)| ns == namespace && m == mode && *b >= bucket - 1)
.fold((0usize, 0i64, 0i64), |(count, sum, max), (_, (c, s, m))| {
(count + c, sum + s, max.max(*m))
});
Ok(JobIndexWaitStats {
count,
avg_ms: if count == 0 { 0 } else { sum_ms / count as i64 },
max_ms,
})
}
async fn job_index_sweep_expired(
&self,
namespace: &str,
now: UtcDateTime,
limit: usize,
) -> anyhow::Result<usize> {
let mut state = self.state.lock().await;
let expired: Vec<(String, String)> = state
.job_index
.iter()
.filter(|((ns, _), row)| ns == namespace && is_expired(row, now))
.take(limit)
.map(|(key, _)| key.clone())
.collect();
let removed = expired.len();
for key in expired {
state.job_index.remove(&key);
}
let stale_bucket = minute_bucket(now) - 5;
state
.transition_metrics
.retain(|(_, _, bucket), _| *bucket >= stale_bucket);
state
.wait_metrics
.retain(|(_, _, bucket), _| *bucket >= stale_bucket);
Ok(removed)
}
}
fn is_expired(row: &JobIndexRow, now: UtcDateTime) -> bool {
row.date_expire.is_some_and(|expire| expire <= now)
}