use crate::{
ClientError, Error, Result,
cache::Cache,
client::Client,
commands::{
ClientTrackingOptions, ClientTrackingStatus, ClusterCommands, ConnectionCommands,
FlushingMode, HashCommands, ServerCommands, StringCommands,
},
network::{sleep, timeout},
resp::cmd,
tests::{get_cluster_test_client, get_default_config, get_test_client, log_try_init},
};
use serial_test::serial;
use std::time::Duration;
#[tokio::test]
#[serial]
async fn cache_get() -> Result<()> {
log_try_init();
let client1 = Client::connect("redis://127.0.0.1?connection_name=client1").await?;
let client2 = Client::connect("redis://127.0.0.1?connection_name=client2").await?;
client2.flushall(FlushingMode::Sync).await?;
client1
.client_tracking(ClientTrackingStatus::Off, ClientTrackingOptions::default())
.await?;
client2.set("key", "value").await?;
let cache = Cache::new(client1.clone(), 60, ClientTrackingOptions::default()).await?;
let value: String = cache.get("key").await?;
assert_eq!("value", value);
let value: String = cache.get("key").await?;
assert_eq!("value", value);
client2.set("key", "new_value").await?;
sleep(Duration::from_millis(100)).await;
let value: String = cache.get("key").await?;
assert_eq!("new_value", value);
let value: String = cache.get("key").await?;
assert_eq!("new_value", value);
Ok(())
}
#[tokio::test]
#[serial]
async fn cache_key_serializing_to_zero_args_errors_instead_of_panicking() -> Result<()> {
log_try_init();
let client = Client::connect("redis://127.0.0.1?connection_name=cache_bad_key").await?;
let cache = Cache::new(client, 60, ClientTrackingOptions::default()).await?;
let result: Result<String> = cache.get(None::<String>).await;
assert!(
matches!(result, Err(Error::Client(ClientError::InvalidCacheKey))),
"expected InvalidCacheKey, got {result:?}"
);
Ok(())
}
#[tokio::test]
#[serial]
async fn cache_hash() -> Result<()> {
log_try_init();
let client1 = Client::connect("redis://127.0.0.1?connection_name=client1").await?;
let client2 = Client::connect("redis://127.0.0.1?connection_name=client2").await?;
client2.flushall(FlushingMode::Sync).await?;
client1
.client_tracking(ClientTrackingStatus::Off, ClientTrackingOptions::default())
.await?;
client2
.hset("key", [("field1", "value1"), ("field2", "value2")])
.await?;
let cache = Cache::new(client1.clone(), 60, ClientTrackingOptions::default()).await?;
let mut values: Vec<(String, String)> = cache.hgetall("key").await?;
values.sort_by(|(f1, _), (f2, _)| f1.cmp(f2));
assert_eq!(
vec![
("field1".to_string(), "value1".to_string()),
("field2".to_string(), "value2".to_string())
],
values
);
let mut values: Vec<(String, String)> = cache.hgetall("key").await?;
values.sort_by(|(f1, _), (f2, _)| f1.cmp(f2));
assert_eq!(
vec![
("field1".to_string(), "value1".to_string()),
("field2".to_string(), "value2".to_string())
],
values
);
let len = cache.hlen("key").await?;
assert_eq!(2, len);
let len = cache.hlen("key").await?;
assert_eq!(2, len);
client2
.hset("key", [("field1", "value11"), ("field2", "value22")])
.await?;
sleep(Duration::from_millis(100)).await;
let mut values: Vec<(String, String)> = cache.hgetall("key").await?;
values.sort_by(|(f1, _), (f2, _)| f1.cmp(f2));
assert_eq!(
vec![
("field1".to_string(), "value11".to_string()),
("field2".to_string(), "value22".to_string())
],
values
);
let mut values: Vec<(String, String)> = cache.hgetall("key").await?;
values.sort_by(|(f1, _), (f2, _)| f1.cmp(f2));
assert_eq!(
vec![
("field1".to_string(), "value11".to_string()),
("field2".to_string(), "value22".to_string())
],
values
);
let len = cache.hlen("key").await?;
assert_eq!(2, len);
let len = cache.hlen("key").await?;
assert_eq!(2, len);
Ok(())
}
#[tokio::test]
#[serial]
async fn cache_mget() -> Result<()> {
log_try_init();
let client1 = Client::connect("redis://127.0.0.1?connection_name=client1").await?;
let client2 = Client::connect("redis://127.0.0.1?connection_name=client2").await?;
client2.flushall(FlushingMode::Sync).await?;
client1
.client_tracking(ClientTrackingStatus::Off, ClientTrackingOptions::default())
.await?;
let cache = Cache::new(client1.clone(), 60, ClientTrackingOptions::default()).await?;
client2
.mset([("key1", "value1"), ("key2", "value2")])
.await?;
assert_eq!("value1", cache.get::<String>("key1").await?);
let values: Vec<String> = cache.mget(["key1", "key2"]).await?;
assert_eq!(vec!["value1".to_string(), "value2".to_string()], values);
assert_eq!("value1", cache.get::<String>("key1").await?);
assert_eq!("value2", cache.get::<String>("key2").await?);
let values: Vec<String> = cache.mget(["key1", "key2"]).await?;
assert_eq!(vec!["value1".to_string(), "value2".to_string()], values);
Ok(())
}
#[tokio::test]
#[serial]
async fn cache_survives_reconnection() -> Result<()> {
log_try_init();
let client1 = Client::connect("redis://127.0.0.1?connection_name=client1").await?;
let client2 = Client::connect("redis://127.0.0.1?connection_name=client2").await?;
client2.flushall(FlushingMode::Sync).await?;
client1
.client_tracking(ClientTrackingStatus::Off, ClientTrackingOptions::default())
.await?;
client2.set("key", "value").await?;
let cache = Cache::new(client1.clone(), 60, ClientTrackingOptions::default()).await?;
assert_eq!("value", cache.get::<String>("key").await?);
client1.send_and_forget(cmd("PING").kill_connection_on_read(1), None)?;
client2.set("key", "new_value").await?;
sleep(Duration::from_millis(2500)).await;
assert_eq!(
"new_value",
cache.get::<String>("key").await?,
"a value written during the outage must not be served from the cache"
);
client2.set("key", "newer_value").await?;
sleep(Duration::from_millis(100)).await;
assert_eq!(
"newer_value",
cache.get::<String>("key").await?,
"invalidations must still be delivered after a reconnection"
);
Ok(())
}
#[tokio::test]
#[serial]
async fn cluster_cache_is_invalidated_for_keys_on_every_shard() -> Result<()> {
log_try_init();
let cached_client = get_cluster_test_client().await?;
let writer = get_cluster_test_client().await?;
writer.flushall(FlushingMode::Sync).await?;
for key in ["key1", "key2", "key3"] {
writer.set(key, "before").await?;
}
let mut slots = Vec::new();
for key in ["key1", "key2", "key3"] {
slots.push(writer.cluster_keyslot(key).await?);
}
assert!(
slots.iter().any(|s| *s > 5460),
"the probe keys must span more than one shard, slots: {slots:?}"
);
let cache = Cache::new(cached_client.clone(), 60, ClientTrackingOptions::default()).await?;
for key in ["key1", "key2", "key3"] {
assert_eq!("before", cache.get::<String>(key).await?);
}
for key in ["key1", "key2", "key3"] {
writer.set(key, "after").await?;
}
sleep(Duration::from_millis(200)).await;
for key in ["key1", "key2", "key3"] {
assert_eq!(
"after",
cache.get::<String>(key).await?,
"the cache must be invalidated for `{key}`, whichever shard holds it"
);
}
Ok(())
}
#[tokio::test]
#[serial]
async fn a_lost_invalidation_flushes_the_cache_instead_of_serving_stale_data() -> Result<()> {
log_try_init();
const PREFIX: &str = "lost_invalidation_key_";
const OTHER_KEYS: usize = 5_000;
let mut config = get_default_config()?;
config.backpressure.max_push_bytes = 1;
let cached_client = Client::connect(config).await?;
let writer = get_test_client().await?;
let target = format!("{PREFIX}target");
writer.set(&target, "v1").await?;
let cache = Cache::new(
cached_client.clone(),
60,
ClientTrackingOptions::default()
.prefix(PREFIX)
.broadcasting(),
)
.await?;
let value: String = cache.get(&target).await?;
assert_eq!("v1", value, "the value under test must start out cached");
writer.send_and_forget(cmd("SET").arg(&target).arg("v2"), None)?;
for i in 0..OTHER_KEYS {
writer.send_and_forget(cmd("SET").arg(format!("{PREFIX}{i}")).arg("v"), None)?;
}
let _: String = writer.send(cmd("PING"), None).await?;
timeout(Duration::from_secs(30), async {
while cache.flush_generation() == 0 {
tokio::task::yield_now().await;
}
Ok::<(), Error>(())
})
.await??;
let value: String = cache.get(&target).await?;
assert_eq!(
"v2", value,
"a dropped invalidation must have flushed the cache, not left a stale value"
);
cached_client
.client_tracking(ClientTrackingStatus::Off, ClientTrackingOptions::default())
.await?;
Ok(())
}