use std::fmt;
use std::time::Duration;
use redis::aio::{ConnectionLike, ConnectionManager};
use redis::Script;
use crate::gcra::{RateLimitInfo, RateLimited};
use crate::quota::{Nanos, Quota};
use crate::storage::{Storage, StorageError, StorageFuture, StorageKey};
const GCRA_SCRIPT: &str = include_str!("gcra.lua");
const DEFAULT_TIMEOUT: Duration = Duration::from_millis(100);
type Reply = (i64, i64, i64, i64);
pub struct RedisStorage<C = ConnectionManager> {
conn: C,
script: Script,
key_prefix: String,
timeout: Duration,
hash_user_ids: bool,
key_secret: Option<Vec<u8>>,
client_clock: bool,
}
impl<C> RedisStorage<C> {
pub fn new(conn: C) -> Self {
Self {
conn,
script: Script::new(GCRA_SCRIPT),
key_prefix: "trt:".to_owned(),
timeout: DEFAULT_TIMEOUT,
hash_user_ids: true,
key_secret: None,
client_clock: false,
}
}
pub fn key_prefix(mut self, prefix: impl Into<String>) -> Self {
self.key_prefix = prefix.into();
self
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
pub fn plain_user_ids(mut self) -> Self {
self.hash_user_ids = false;
self
}
pub fn key_secret(mut self, secret: impl AsRef<[u8]>) -> Self {
self.key_secret = Some(secret.as_ref().to_vec());
self
}
pub fn use_client_clock(mut self) -> Self {
self.client_clock = true;
self
}
pub fn redis_key(&self, key: StorageKey<'_>) -> String {
if self.hash_user_ids {
let digest = match &self.key_secret {
Some(secret) => hmac_sha1(secret, key.user_id.as_bytes()),
None => sha1_smol::Sha1::from(key.user_id).digest(),
};
format!("{}{}:{}", self.key_prefix, key.tier, digest)
} else {
format!(
"{}{}:{}:{}",
self.key_prefix,
key.tier.len(),
key.tier,
key.user_id
)
}
}
}
#[derive(Debug)]
struct ScriptCall {
key: String,
now: String,
emission_interval: u64,
burst_offset: u64,
cost: u32,
}
impl<C> RedisStorage<C> {
fn script_call(&self, key: StorageKey<'_>, quota: &Quota, cost: u32, now: Nanos) -> ScriptCall {
let emission_interval = micros_ceil(quota.emission_interval_nanos()).max(1);
ScriptCall {
key: self.redis_key(key),
now: if self.client_clock {
(now / 1_000).to_string()
} else {
String::new()
},
emission_interval,
burst_offset: emission_interval.saturating_mul(u64::from(quota.max_burst())),
cost,
}
}
}
impl<C> fmt::Debug for RedisStorage<C> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RedisStorage")
.field("key_prefix", &self.key_prefix)
.field("timeout", &self.timeout)
.field("hash_user_ids", &self.hash_user_ids)
.field("key_secret", &self.key_secret.is_some())
.field("client_clock", &self.client_clock)
.finish_non_exhaustive()
}
}
impl<C> Storage for RedisStorage<C>
where
C: ConnectionLike + Clone + Send + Sync + 'static,
{
fn check_and_update<'a>(
&'a self,
key: StorageKey<'a>,
quota: &'a Quota,
cost: u32,
now: Nanos,
) -> StorageFuture<'a> {
Box::pin(async move {
let call = self.script_call(key, quota, cost, now);
let mut invocation = self.script.prepare_invoke();
invocation
.key(call.key)
.arg(call.now)
.arg(call.emission_interval)
.arg(call.burst_offset)
.arg(call.cost);
let mut conn = self.conn.clone();
let call = invocation.invoke_async::<Reply>(&mut conn);
match tokio::time::timeout(self.timeout, call).await {
Ok(Ok(reply)) => Ok(decode_reply(reply, quota.max_burst())),
Ok(Err(err)) => Err(StorageError(Box::new(err))),
Err(_) => Err(StorageError(
format!("redis did not answer within {:?}", self.timeout).into(),
)),
}
})
}
}
fn hmac_sha1(secret: &[u8], message: &[u8]) -> sha1_smol::Digest {
const BLOCK: usize = 64;
let mut key = [0u8; BLOCK];
if secret.len() > BLOCK {
key[..20].copy_from_slice(&sha1_smol::Sha1::from(secret).digest().bytes());
} else {
key[..secret.len()].copy_from_slice(secret);
}
let mut inner = sha1_smol::Sha1::new();
inner.update(&key.map(|b| b ^ 0x36));
inner.update(message);
let mut outer = sha1_smol::Sha1::new();
outer.update(&key.map(|b| b ^ 0x5c));
outer.update(&inner.digest().bytes());
outer.digest()
}
fn micros_ceil(nanos: Nanos) -> u64 {
nanos / 1_000 + u64::from(nanos % 1_000 != 0)
}
fn decode_reply(
(allowed, remaining, retry_after_us, reset_after_us): Reply,
limit: u32,
) -> Result<RateLimitInfo, RateLimited> {
let micros = |value: i64| Duration::from_micros(u64::try_from(value).unwrap_or(0));
if allowed == 1 {
Ok(RateLimitInfo {
limit,
remaining: u32::try_from(remaining.max(0)).unwrap_or(u32::MAX),
reset_after: micros(reset_after_us),
})
} else {
Err(RateLimited {
limit,
retry_after: micros(retry_after_us),
reset_after: micros(reset_after_us),
})
}
}
#[cfg(test)]
mod script_tests;
#[cfg(test)]
mod tests {
use redis::{Cmd, Pipeline, RedisFuture, Value};
use super::*;
#[derive(Clone)]
struct NeverAnswers;
impl ConnectionLike for NeverAnswers {
fn req_packed_command<'a>(&'a mut self, _cmd: &'a Cmd) -> RedisFuture<'a, Value> {
Box::pin(std::future::pending())
}
fn req_packed_commands<'a>(
&'a mut self,
_pipeline: &'a Pipeline,
_offset: usize,
_count: usize,
) -> RedisFuture<'a, Vec<Value>> {
Box::pin(std::future::pending())
}
fn get_db(&self) -> i64 {
0
}
}
#[derive(Clone)]
struct Refused;
impl ConnectionLike for Refused {
fn req_packed_command<'a>(&'a mut self, _cmd: &'a Cmd) -> RedisFuture<'a, Value> {
Box::pin(std::future::ready(Err(refused())))
}
fn req_packed_commands<'a>(
&'a mut self,
_pipeline: &'a Pipeline,
_offset: usize,
_count: usize,
) -> RedisFuture<'a, Vec<Value>> {
Box::pin(std::future::ready(Err(refused())))
}
fn get_db(&self) -> i64 {
0
}
}
fn refused() -> redis::RedisError {
std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "connection refused").into()
}
#[test]
fn keys_hash_user_ids_by_default() {
let storage = RedisStorage::new(NeverAnswers);
let key = storage.redis_key(StorageKey::new("sk_live_123", "pro"));
assert_eq!(
key,
format!("trt:pro:{}", sha1_smol::Sha1::from("sk_live_123").digest())
);
assert!(!key.contains("sk_live_123"));
}
#[test]
fn keys_never_collide_across_separators() {
for storage in [
RedisStorage::new(NeverAnswers),
RedisStorage::new(NeverAnswers).plain_user_ids(),
] {
for (first, second) in [(("a:b", "c"), ("a", "b:c")), (("c", "a:b"), ("b:c", "a"))] {
let a = storage.redis_key(StorageKey::new(first.0, first.1));
let b = storage.redis_key(StorageKey::new(second.0, second.1));
assert_ne!(a, b, "{first:?} and {second:?}");
}
}
}
#[test]
fn plain_keys_keep_the_user_id_readable() {
let storage = RedisStorage::new(NeverAnswers)
.key_prefix("app:")
.plain_user_ids();
let key = storage.redis_key(StorageKey::new("alice", "free"));
assert_eq!(key, "app:4:free:alice");
}
#[test]
fn hmac_sha1_matches_rfc_2202() {
let cases: [(&[u8], &[u8], &str); 3] = [
(
&[0x0b; 20],
b"Hi There",
"b617318655057264e28bc0b6fb378c8ef146be00",
),
(
b"Jefe",
b"what do ya want for nothing?",
"effcdf6ae5eb2fa2d27416d5f184df9c259a7c79",
),
(
&[0xaa; 80],
b"Test Using Larger Than Block-Size Key - Hash Key First",
"aa4ae5e15272d00e95705637ce8a3b55ed402112",
),
];
for (secret, message, expected) in cases {
assert_eq!(hmac_sha1(secret, message).to_string(), expected);
}
}
#[test]
fn a_key_secret_changes_every_key() {
let key = StorageKey::new("203.0.113.7", "free");
let plain_hash = RedisStorage::new(NeverAnswers).redis_key(key);
let first = RedisStorage::new(NeverAnswers)
.key_secret("one")
.redis_key(key);
let again = RedisStorage::new(NeverAnswers)
.key_secret("one")
.redis_key(key);
let second = RedisStorage::new(NeverAnswers)
.key_secret("two")
.redis_key(key);
assert_eq!(first, again, "the same secret must give the same key");
assert_ne!(first, plain_hash);
assert_ne!(first, second);
assert!(first.starts_with("trt:free:") && first.len() == "trt:free:".len() + 40);
assert!(!first.contains("203.0.113.7"));
}
#[test]
fn debug_never_shows_the_key_secret() {
let storage = RedisStorage::new(NeverAnswers).key_secret("hunter2-secret");
let out = format!("{storage:?}");
assert!(!out.contains("hunter2"), "{out}");
assert!(out.contains("key_secret: true"), "{out}");
}
#[test]
fn micros_round_up() {
assert_eq!(micros_ceil(333_333_333), 333_334);
assert_eq!(micros_ceil(250_000_000), 250_000);
assert_eq!(micros_ceil(1), 1);
}
#[test]
fn replies_decode_into_the_gcra_decision() {
let allowed = decode_reply((1, 3, 0, 1_500_000), 4);
assert_eq!(
allowed,
Ok(RateLimitInfo {
limit: 4,
remaining: 3,
reset_after: Duration::from_millis(1_500),
})
);
let limited = decode_reply((0, 0, 250_000, 1_000_000), 4);
assert_eq!(
limited,
Err(RateLimited {
limit: 4,
retry_after: Duration::from_millis(250),
reset_after: Duration::from_secs(1),
})
);
}
#[tokio::test(start_paused = true)]
async fn a_silent_redis_times_out_as_a_storage_error() {
let storage = RedisStorage::new(NeverAnswers).timeout(Duration::from_millis(50));
let quota = Quota::per_second(1);
let result = storage
.check_and_update(StorageKey::new("u1", "free"), "a, 1, 0)
.await;
let err = result.expect_err("a check with no answer must fail");
assert!(
err.to_string().contains("did not answer within 50ms"),
"{err}"
);
}
#[tokio::test]
async fn connection_errors_become_storage_errors() {
let storage = RedisStorage::new(Refused);
let quota = Quota::per_second(1);
let result = storage
.check_and_update(StorageKey::new("u1", "free"), "a, 1, 0)
.await;
let err = result.expect_err("a refused connection must fail");
assert!(err.to_string().contains("connection refused"), "{err}");
}
}