use reqwest_drive::{
CacheBust, CacheBypass, CachePolicy, ThrottlePolicy, init_cache, init_cache_process_scoped,
init_cache_process_scoped_with_throttle, init_cache_with_drive,
init_cache_with_drive_and_throttle, init_cache_with_throttle,
init_client_with_cache_and_throttle, init_throttle,
};
use reqwest_middleware::ClientBuilder;
use simd_r_drive::DataStore;
use std::env;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, MutexGuard, OnceLock};
use std::time::Duration;
use tempfile::TempDir;
use tokio::sync::{Barrier, mpsc};
use tokio::time::{Instant, sleep};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn cwd_test_lock() -> &'static Mutex<()> {
static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
LOCK.get_or_init(|| Mutex::new(()))
}
struct CwdGuard {
previous: PathBuf,
_cwd_lock: MutexGuard<'static, ()>,
}
impl CwdGuard {
fn swap_to(path: &Path) -> std::io::Result<Self> {
let cwd_lock_guard = cwd_test_lock().lock().expect("acquire cwd test lock");
let previous = env::current_dir()?;
env::set_current_dir(path)?;
Ok(Self {
previous,
_cwd_lock: cwd_lock_guard,
})
}
}
impl Drop for CwdGuard {
fn drop(&mut self) {
let _ = env::set_current_dir(&self.previous);
}
}
#[tokio::test]
async fn test_cache_middleware() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache.bin");
let mock_server = MockServer::start().await;
let template = ResponseTemplate::new(200)
.set_body_string("cached response")
.insert_header("Cache-Control", "max-age=60");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(template.clone())
.mount(&mock_server)
.await;
let cache = init_cache(&cache_path, CachePolicy::default());
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache.clone())
.build();
let url = mock_server.uri();
let first_response = client.get(&url).send().await.unwrap();
let first_body = first_response.text().await.unwrap();
assert_eq!(first_body, "cached response");
let second_response = client.get(&url).send().await.unwrap();
let second_body = second_response.text().await.unwrap();
assert_eq!(second_body, "cached response");
}
#[tokio::test]
async fn test_init_cache_process_scoped() {
let temp_root = TempDir::new().expect("create tempfile root");
let _cwd_guard = CwdGuard::swap_to(temp_root.path()).expect("set cwd to tempfile root");
let mock_server = MockServer::start().await;
Mock::given(wiremock::matchers::method("GET"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string("process scoped cache")
.insert_header("Cache-Control", "max-age=60"),
)
.expect(1)
.mount(&mock_server)
.await;
let cache = init_cache_process_scoped(CachePolicy::default())
.expect("failed to initialize process-scoped cache");
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache)
.build();
let url = mock_server.uri();
let first = client.get(&url).send().await.unwrap().text().await.unwrap();
let second = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(first, "process scoped cache");
assert_eq!(second, "process scoped cache");
let discovered_group = temp_root.path().join(".cache").join(env!("CARGO_PKG_NAME"));
assert!(discovered_group.exists());
}
#[tokio::test]
async fn test_init_cache_process_scoped_with_throttle() {
let temp_root = TempDir::new().expect("create tempfile root");
let _cwd_guard = CwdGuard::swap_to(temp_root.path()).expect("set cwd to tempfile root");
let mock_server = MockServer::start().await;
Mock::given(wiremock::matchers::method("GET"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string("process scoped cache + throttle")
.insert_header("Cache-Control", "max-age=60"),
)
.expect(1)
.mount(&mock_server)
.await;
let (cache, throttle) = init_cache_process_scoped_with_throttle(
CachePolicy::default(),
ThrottlePolicy {
base_delay_ms: 50,
adaptive_jitter_ms: 0,
max_concurrent: 1,
max_retries: 0,
},
)
.expect("failed to initialize process-scoped cache with throttle");
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache)
.with_arc(throttle)
.build();
let url = mock_server.uri();
let first = client.get(&url).send().await.unwrap().text().await.unwrap();
let second = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(first, "process scoped cache + throttle");
assert_eq!(second, "process scoped cache + throttle");
let discovered_group = temp_root.path().join(".cache").join(env!("CARGO_PKG_NAME"));
assert!(discovered_group.exists());
}
#[tokio::test]
async fn test_cache_key_normalizes_query_param_order() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_query_normalization.bin");
let mock_server = MockServer::start().await;
Mock::given(wiremock::matchers::method("GET"))
.respond_with(ResponseTemplate::new(200).set_body_string("query-normalized"))
.expect(1)
.mount(&mock_server)
.await;
let cache = init_cache(&cache_path, CachePolicy::default());
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache)
.build();
let url_a = format!("{}/query?a=1&b=2", mock_server.uri());
let url_b = format!("{}/query?b=2&a=1", mock_server.uri());
let first = client
.get(&url_a)
.send()
.await
.unwrap()
.text()
.await
.unwrap();
let second = client
.get(&url_b)
.send()
.await
.unwrap()
.text()
.await
.unwrap();
assert_eq!(first, "query-normalized");
assert_eq!(second, "query-normalized");
}
#[tokio::test]
async fn test_cache_key_varies_on_accept_language_header() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_vary_header.bin");
let mock_server = MockServer::start().await;
let request_counter = Arc::new(AtomicUsize::new(0));
let counter_clone = Arc::clone(&request_counter);
Mock::given(wiremock::matchers::method("GET"))
.respond_with(move |req: &wiremock::Request| {
let _ = counter_clone.fetch_add(1, Ordering::SeqCst);
let lang = req
.headers
.get("accept-language")
.map(|v| String::from_utf8_lossy(v.as_bytes()).to_string())
.unwrap_or_else(|| "none".to_string());
ResponseTemplate::new(200)
.set_body_string(format!("lang={}", lang))
.insert_header("Cache-Control", "max-age=60")
})
.expect(2)
.mount(&mock_server)
.await;
let cache = init_cache(&cache_path, CachePolicy::default());
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache)
.build();
let url = format!("{}/vary", mock_server.uri());
let en = client
.get(&url)
.header("accept-language", "en-US")
.send()
.await
.unwrap()
.text()
.await
.unwrap();
let fr = client
.get(&url)
.header("accept-language", "fr-FR")
.send()
.await
.unwrap()
.text()
.await
.unwrap();
let en_cached = client
.get(&url)
.header("accept-language", "en-US")
.send()
.await
.unwrap()
.text()
.await
.unwrap();
assert_eq!(en, "lang=en-US");
assert_eq!(fr, "lang=fr-FR");
assert_eq!(en_cached, "lang=en-US");
assert_eq!(request_counter.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_throttling_behavior() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache.bin");
let mock_server = MockServer::start().await;
let response_template = ResponseTemplate::new(200).set_body_string("throttled response");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(response_template.clone())
.expect(3) .mount(&mock_server)
.await;
let cache_policy = CachePolicy::default();
let throttle_policy = ThrottlePolicy {
base_delay_ms: 200, adaptive_jitter_ms: 100, max_concurrent: 1, max_retries: 0, };
let (_, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle.clone())
.build();
let mut timestamps: Vec<std::time::Instant> = vec![];
let start_time = std::time::Instant::now();
for i in 0..3 {
let request_start = std::time::Instant::now(); let url = format!("{}/test{}", mock_server.uri(), i); let response = client.get(&url).send().await.unwrap();
let body = response.text().await.unwrap();
assert_eq!(body, "throttled response");
timestamps.push(request_start);
}
let elapsed = start_time.elapsed();
let min_expected_delay = Duration::from_millis(400); let max_expected_delay = Duration::from_millis(800);
if elapsed < min_expected_delay {
panic!(
"Throttling was too fast! Expected at least {:?}, but got {:?}",
min_expected_delay, elapsed
);
} else if elapsed > max_expected_delay {
tracing::warn!(
"⚠️ Warning: Throttling took longer than expected. Expected at most {:?}, but got {:?}",
max_expected_delay,
elapsed
);
}
for window in timestamps.windows(2) {
let delay_between_requests = window[1].duration_since(window[0]); let min_per_request = Duration::from_millis(200);
let max_per_request = Duration::from_millis(400);
assert!(
delay_between_requests >= min_per_request,
"Request spacing too short: {:?}",
delay_between_requests
);
if delay_between_requests > max_per_request {
tracing::warn!(
"⚠️ Warning: Request spacing exceeded max expected ({:?}). Got {:?}",
max_per_request,
delay_between_requests
);
}
}
}
#[tokio::test]
async fn test_cache_expiration() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache.bin");
let mock_server = MockServer::start().await;
let template = ResponseTemplate::new(200)
.set_body_string("temporary cache")
.insert_header("Cache-Control", "max-age=1");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(template)
.mount(&mock_server)
.await;
let cache_policy = CachePolicy {
default_ttl: Duration::from_secs(1), respect_headers: true,
cache_status_override: None,
};
let cache = init_cache(&cache_path, cache_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache.clone())
.build();
let url = mock_server.uri();
let first_response = client.get(&url).send().await.unwrap();
let first_body = first_response.text().await.unwrap();
assert_eq!(first_body, "temporary cache");
sleep(Duration::from_secs(2)).await;
let second_response = client.get(&url).send().await.unwrap();
let second_body = second_response.text().await.unwrap();
assert_eq!(second_body, "temporary cache");
}
#[tokio::test]
async fn test_backoff_on_server_error() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache.bin");
let mock_server = MockServer::start().await;
let error_template = ResponseTemplate::new(500);
Mock::given(wiremock::matchers::method("GET"))
.respond_with(error_template.clone())
.expect(3) .mount(&mock_server)
.await;
let cache_policy = CachePolicy::default();
let throttle_policy = ThrottlePolicy {
base_delay_ms: 100, adaptive_jitter_ms: 50, max_concurrent: 1, max_retries: 2, };
let (_, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle.clone())
.build();
let url = format!("{}/test-server-error", mock_server.uri());
let start_time = std::time::Instant::now();
let response = client.get(&url).send().await.unwrap();
let elapsed = start_time.elapsed();
assert_eq!(
response.status(),
500,
"Expected final response to be HTTP 500, but got {:?}",
response.status()
);
let min_expected_delay = Duration::from_millis(200); let max_expected_delay = Duration::from_millis(800);
assert!(
elapsed >= min_expected_delay,
"Backoff was too fast! Expected at least {:?}, but got {:?}",
min_expected_delay,
elapsed
);
if elapsed > max_expected_delay {
tracing::warn!(
"⚠️ Warning: Backoff took longer than expected. Expected at most {:?}, but got {:?}",
max_expected_delay,
elapsed
);
}
}
#[tokio::test]
async fn test_backoff_with_eventual_success() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache.bin");
let mock_server = MockServer::start().await;
let error_template = ResponseTemplate::new(500); let success_template = ResponseTemplate::new(200).set_body_string("success");
let request_counter = Arc::new(AtomicUsize::new(0));
let counter_clone = request_counter.clone();
Mock::given(wiremock::matchers::method("GET"))
.respond_with(move |_: &wiremock::Request| {
let count = counter_clone.fetch_add(1, Ordering::SeqCst);
if count < 2 {
error_template.clone() } else {
success_template.clone() }
})
.expect(3) .mount(&mock_server)
.await;
let cache_policy = CachePolicy::default();
let throttle_policy = ThrottlePolicy {
base_delay_ms: 100, adaptive_jitter_ms: 50, max_concurrent: 1, max_retries: 2, };
let (_, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle.clone())
.build();
let url = format!("{}/test-backoff-retry", mock_server.uri());
let start_time = std::time::Instant::now();
let response = client.get(&url).send().await.unwrap();
let body = response.text().await.unwrap();
let elapsed = start_time.elapsed();
assert_eq!(
body, "success",
"Expected final response to be 'success', but got {:?}",
body
);
let min_expected_delay = Duration::from_millis(500); let max_expected_delay = Duration::from_millis(1200);
assert!(
elapsed >= min_expected_delay,
"Backoff was too fast! Expected at least {:?}, but got {:?}",
min_expected_delay,
elapsed
);
if elapsed > max_expected_delay {
tracing::warn!(
"⚠️ Warning: Backoff took longer than expected. Expected at most {:?}, but got {:?}",
max_expected_delay,
elapsed
);
}
assert_eq!(
request_counter.load(Ordering::SeqCst),
3,
"Expected exactly 3 requests (2 failures + 1 success), but got {}",
request_counter.load(Ordering::SeqCst)
);
}
#[tokio::test]
async fn test_init_cache_with_throttle() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_throttle.bin");
let mock_server = MockServer::start().await;
let response_template = ResponseTemplate::new(200)
.set_body_string("cached response")
.insert_header("Cache-Control", "max-age=60");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(response_template.clone())
.expect(1) .mount(&mock_server)
.await;
let cache_policy = CachePolicy {
default_ttl: Duration::from_secs(60), respect_headers: true,
cache_status_override: None,
};
let throttle_policy = ThrottlePolicy {
base_delay_ms: 200, adaptive_jitter_ms: 50, max_concurrent: 1, max_retries: 1, };
let (cache, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache.clone())
.with_arc(throttle.clone())
.build();
let url = mock_server.uri();
let start_time_1 = std::time::Instant::now();
let first_response = client.get(&url).send().await.unwrap();
let first_body = first_response.text().await.unwrap();
let elapsed_1 = start_time_1.elapsed();
assert_eq!(first_body, "cached response");
assert!(
elapsed_1 >= Duration::from_millis(200),
"First request was too fast!"
);
let start_time_2 = std::time::Instant::now();
let second_response = client.get(&url).send().await.unwrap();
let second_body = second_response.text().await.unwrap();
let elapsed_2 = start_time_2.elapsed();
assert_eq!(second_body, "cached response");
assert!(
elapsed_2 < Duration::from_millis(50),
"Second request was not instant despite caching!"
);
tracing::info!(
"Test passed! First request took {:?}, second request took {:?} (should be cached).",
elapsed_1,
elapsed_2
);
}
#[tokio::test]
async fn test_with_drive_arc() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_drive.bin");
let mock_server = MockServer::start().await;
let response_template = ResponseTemplate::new(200)
.set_body_string("cached response")
.insert_header("Cache-Control", "max-age=60");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(response_template.clone())
.expect(1) .mount(&mock_server)
.await;
let store = Arc::new(DataStore::open(&cache_path).unwrap());
let cache_policy = CachePolicy::default();
let cache = init_cache_with_drive(store.clone(), cache_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache.clone())
.build();
let url = mock_server.uri();
let first_response = client.get(&url).send().await.unwrap();
let first_body = first_response.text().await.unwrap();
assert_eq!(first_body, "cached response");
let second_response = client.get(&url).send().await.unwrap();
let second_body = second_response.text().await.unwrap();
assert_eq!(second_body, "cached response");
tracing::info!("`init_cache_with_drive` successfully initialized and cached responses.");
}
#[tokio::test]
async fn test_with_drive_arc_and_throttle() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_drive_throttle.bin");
let mock_server = MockServer::start().await;
let response_template = ResponseTemplate::new(200).set_body_string("throttled response");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(response_template.clone())
.expect(3) .mount(&mock_server)
.await;
let store = Arc::new(DataStore::open(&cache_path).unwrap());
let cache_policy = CachePolicy::default();
let throttle_policy = ThrottlePolicy {
base_delay_ms: 200, adaptive_jitter_ms: 100, max_concurrent: 1, max_retries: 0, };
let (cache, throttle) =
init_cache_with_drive_and_throttle(store.clone(), cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle.clone()) .with_arc(cache.clone()) .build();
let mut timestamps: Vec<std::time::Instant> = vec![];
let start_time = std::time::Instant::now();
for i in 0..3 {
let request_start = std::time::Instant::now(); let url = format!("{}/test{}", mock_server.uri(), i); let response = client.get(&url).send().await.unwrap();
let body = response.text().await.unwrap();
assert_eq!(body, "throttled response");
timestamps.push(request_start);
}
let elapsed = start_time.elapsed();
let min_expected_delay = Duration::from_millis(400); let max_expected_delay = Duration::from_millis(800);
if elapsed < min_expected_delay {
panic!(
"Throttling was too fast! Expected at least {:?}, but got {:?}",
min_expected_delay, elapsed
);
} else if elapsed > max_expected_delay {
tracing::warn!(
"⚠️ Warning: Throttling took longer than expected. Expected at most {:?}, but got {:?}",
max_expected_delay,
elapsed
);
}
for window in timestamps.windows(2) {
let delay_between_requests = window[1].duration_since(window[0]); let min_per_request = Duration::from_millis(200);
let max_per_request = Duration::from_millis(400);
assert!(
delay_between_requests >= min_per_request,
"Request spacing too short: {:?}",
delay_between_requests
);
if delay_between_requests > max_per_request {
tracing::warn!(
"⚠️ Warning: Request spacing exceeded max expected ({:?}). Got {:?}",
max_per_request,
delay_between_requests
);
}
}
tracing::info!("`init_cache_with_drive_and_throttle` successfully enforced throttling.");
}
#[tokio::test]
async fn test_cache_status_override() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_override.bin");
let mock_server = MockServer::start().await;
let success_template = ResponseTemplate::new(200).set_body_string("success response");
let not_found_template = ResponseTemplate::new(404).set_body_string("not found response");
let server_error_template = ResponseTemplate::new(500).set_body_string("server error response");
Mock::given(wiremock::matchers::path("/success"))
.respond_with(success_template.clone())
.expect(1)
.mount(&mock_server)
.await;
Mock::given(wiremock::matchers::path("/not_found"))
.respond_with(not_found_template.clone())
.expect(1)
.mount(&mock_server)
.await;
Mock::given(wiremock::matchers::path("/server_error"))
.respond_with(server_error_template.clone())
.expect(2)
.mount(&mock_server)
.await;
let cache_policy = CachePolicy {
default_ttl: Duration::from_secs(60),
respect_headers: true,
cache_status_override: Some(vec![200, 404]), };
let cache = init_cache(&cache_path, cache_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache.clone())
.build();
let success_url = format!("{}/success", mock_server.uri());
let not_found_url = format!("{}/not_found", mock_server.uri());
let server_error_url = format!("{}/server_error", mock_server.uri());
let first_success_response = client.get(&success_url).send().await.unwrap();
let first_success_body = first_success_response.text().await.unwrap();
assert_eq!(first_success_body, "success response");
let first_not_found_response = client.get(¬_found_url).send().await.unwrap();
let first_not_found_body = first_not_found_response.text().await.unwrap();
assert_eq!(first_not_found_body, "not found response");
let first_server_error_response = client.get(&server_error_url).send().await.unwrap();
let first_server_error_body = first_server_error_response.text().await.unwrap();
assert_eq!(first_server_error_body, "server error response");
let second_success_response = client.get(&success_url).send().await.unwrap();
let second_success_body = second_success_response.text().await.unwrap();
assert_eq!(second_success_body, "success response");
let second_not_found_response = client.get(¬_found_url).send().await.unwrap();
let second_not_found_body = second_not_found_response.text().await.unwrap();
assert_eq!(second_not_found_body, "not found response");
let second_server_error_response = client.get(&server_error_url).send().await.unwrap();
let second_server_error_body = second_server_error_response.text().await.unwrap();
assert_eq!(second_server_error_body, "server error response");
}
#[tokio::test]
async fn test_concurrent_requests_without_cache() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_concurrent.bin");
let mock_server = MockServer::start().await;
let response_template = ResponseTemplate::new(200).set_body_string("concurrent response");
let num_requests = 5;
let max_concurrent = 3;
for i in 0..num_requests {
let path = format!("/test{}", i);
Mock::given(wiremock::matchers::path(path.clone()))
.respond_with(response_template.clone())
.expect(1) .mount(&mock_server)
.await;
}
let cache_policy = CachePolicy {
default_ttl: Duration::from_secs(60),
respect_headers: true,
cache_status_override: None,
};
let throttle_policy = ThrottlePolicy {
base_delay_ms: 0, adaptive_jitter_ms: 0,
max_concurrent,
max_retries: 0,
};
let (_, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle.clone())
.build();
let barrier = Arc::new(Barrier::new(num_requests));
let (tx, mut rx) = mpsc::channel(num_requests);
let handles: Vec<_> = (0..num_requests)
.map(|i| {
let client = client.clone();
let url = format!("{}/test{}", mock_server.uri(), i);
let barrier = barrier.clone();
let tx = tx.clone();
tokio::spawn(async move {
barrier.wait().await;
let start_time = Instant::now();
let response = client.get(&url).send().await.unwrap();
let body = response.text().await.unwrap();
assert_eq!(body, "concurrent response");
let elapsed_time = start_time.elapsed();
tx.send(elapsed_time).await.unwrap(); })
})
.collect();
for handle in handles {
handle.await.unwrap();
}
drop(tx);
let mut timestamps: Vec<Duration> = vec![];
while let Some(elapsed) = rx.recv().await {
timestamps.push(elapsed);
}
timestamps.sort();
let first_batch = timestamps
.iter()
.take(max_concurrent)
.cloned()
.collect::<Vec<_>>();
assert!(
first_batch.len() == max_concurrent,
"Expected {} concurrent requests but only got {}",
max_concurrent,
first_batch.len()
);
}
#[tokio::test]
async fn test_throttling_respects_max_concurrent() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_throttle_test.bin");
let mock_server = MockServer::start().await;
let response_delay = Duration::from_millis(250); let response_template = ResponseTemplate::new(200)
.set_body_string("throttled response")
.set_delay(response_delay);
let num_requests = 10;
let max_concurrent = 3;
for i in 0..num_requests {
let path = format!("/test{}", i);
Mock::given(wiremock::matchers::path(path.clone()))
.respond_with(response_template.clone())
.expect(1) .mount(&mock_server)
.await;
}
let cache_policy = CachePolicy {
default_ttl: Duration::from_secs(60),
respect_headers: true,
cache_status_override: None,
};
let throttle_policy = ThrottlePolicy {
base_delay_ms: 0, adaptive_jitter_ms: 0,
max_concurrent,
max_retries: 0,
};
let (_, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle.clone())
.build();
let barrier = Arc::new(Barrier::new(num_requests));
let max_seen_concurrent = Arc::new(AtomicUsize::new(0));
let handles: Vec<_> = (0..num_requests)
.map(|i| {
let client = client.clone();
let url = format!("{}/test{}", mock_server.uri(), i);
let barrier = barrier.clone();
let max_seen = Arc::clone(&max_seen_concurrent);
let throttle = throttle.clone();
tokio::spawn(async move {
barrier.wait().await;
let start_time = Instant::now();
tracing::debug!("Awaiting the lock...");
let current_in_flight = max_concurrent - throttle.available_permits();
tracing::debug!("Current in flight: {}", current_in_flight);
max_seen.fetch_max(current_in_flight, Ordering::SeqCst);
let response = client.get(&url).send().await.unwrap();
let body = response.text().await.unwrap();
assert_eq!(body, "throttled response");
start_time.elapsed()
})
})
.collect();
for handle in handles {
handle.await.unwrap();
}
let max_seen = max_seen_concurrent.load(Ordering::SeqCst);
assert!(
max_seen <= max_concurrent,
"Max concurrent requests exceeded! Expected at most {}, but saw {}.",
max_concurrent,
max_seen
);
tracing::info!(
"✅ Throttling enforced correctly! Max concurrent requests: {} (expected {}).",
max_seen,
max_concurrent
);
}
#[tokio::test]
async fn test_throttle_policy_override() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_throttle_override.bin");
let mock_server = MockServer::start().await;
let error_template = ResponseTemplate::new(500); let success_template = ResponseTemplate::new(200).set_body_string("eventual success");
let request_counter = Arc::new(AtomicUsize::new(0));
let counter_clone = request_counter.clone();
Mock::given(wiremock::matchers::method("GET"))
.respond_with(move |_: &wiremock::Request| {
let count = counter_clone.fetch_add(1, Ordering::SeqCst);
if count < 1 {
error_template.clone() } else {
success_template.clone() }
})
.expect(2) .mount(&mock_server)
.await;
let cache_policy = CachePolicy::default();
let throttle_policy = ThrottlePolicy {
base_delay_ms: 100, adaptive_jitter_ms: 50, max_concurrent: 1, max_retries: 3, };
let (_, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle.clone())
.build();
let url = format!("{}/test-throttle-override", mock_server.uri());
let start_time = std::time::Instant::now();
let custom_throttle = ThrottlePolicy {
base_delay_ms: 100, adaptive_jitter_ms: 50, max_concurrent: 1, max_retries: 1, };
let mut request = client.get(&url);
request.extensions().insert(custom_throttle); let response = request.send().await.unwrap();
let body = response.text().await.unwrap();
let elapsed = start_time.elapsed();
assert_eq!(
body, "eventual success",
"Expected final response to be 'eventual success', but got {:?}",
body
);
assert_eq!(
request_counter.load(Ordering::SeqCst),
2,
"Expected exactly 2 requests (1 failure + 1 success), but got {}",
request_counter.load(Ordering::SeqCst)
);
let min_expected_delay = Duration::from_millis(150); let max_expected_delay = Duration::from_millis(400);
assert!(
elapsed >= min_expected_delay,
"Backoff was too fast! Expected at least {:?}, but got {:?}",
min_expected_delay,
elapsed
);
if elapsed > max_expected_delay {
tracing::warn!(
"⚠️ Warning: Backoff took longer than expected. Expected at most {:?}, but got {:?}",
max_expected_delay,
elapsed
);
}
}
#[tokio::test]
async fn test_per_request_cache_bypass_with_throttle_and_shared_store() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_bypass.bin");
let mock_server = MockServer::start().await;
let request_counter = Arc::new(AtomicUsize::new(0));
let counter_clone = Arc::clone(&request_counter);
Mock::given(wiremock::matchers::method("GET"))
.respond_with(move |_: &wiremock::Request| {
let count = counter_clone.fetch_add(1, Ordering::SeqCst);
ResponseTemplate::new(200)
.set_body_string(format!("server-response-{}", count))
.insert_header("Cache-Control", "max-age=60")
})
.expect(2)
.mount(&mock_server)
.await;
let cache_policy = CachePolicy::default();
let throttle_policy = ThrottlePolicy {
base_delay_ms: 100,
adaptive_jitter_ms: 0,
max_concurrent: 1,
max_retries: 0,
};
let (cache, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache)
.with_arc(throttle)
.build();
let url = format!("{}/cache-bypass", mock_server.uri());
let first = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(first, "server-response-0");
let second = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(second, "server-response-0");
let bypass_start = Instant::now();
let mut bypass_request = client.get(&url);
bypass_request.extensions().insert(CacheBypass(true));
let third = bypass_request.send().await.unwrap().text().await.unwrap();
let bypass_elapsed = bypass_start.elapsed();
assert_eq!(third, "server-response-1");
assert!(
bypass_elapsed >= Duration::from_millis(100),
"Bypassed request should still be throttled"
);
let fourth = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(fourth, "server-response-0");
assert_eq!(request_counter.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_per_request_cache_bust_refreshes_cached_value() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("cache_bust.bin");
let mock_server = MockServer::start().await;
let request_counter = Arc::new(AtomicUsize::new(0));
let counter_clone = Arc::clone(&request_counter);
Mock::given(wiremock::matchers::method("GET"))
.respond_with(move |_: &wiremock::Request| {
let count = counter_clone.fetch_add(1, Ordering::SeqCst);
ResponseTemplate::new(200)
.set_body_string(format!("server-response-{}", count))
.insert_header("Cache-Control", "max-age=60")
})
.expect(2)
.mount(&mock_server)
.await;
let cache_policy = CachePolicy::default();
let throttle_policy = ThrottlePolicy {
base_delay_ms: 100,
adaptive_jitter_ms: 0,
max_concurrent: 1,
max_retries: 0,
};
let (cache, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(cache)
.with_arc(throttle)
.build();
let url = format!("{}/cache-bust", mock_server.uri());
let first = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(first, "server-response-0");
let second = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(second, "server-response-0");
let bust_start = Instant::now();
let mut bust_request = client.get(&url);
bust_request.extensions().insert(CacheBust(true));
let third = bust_request.send().await.unwrap().text().await.unwrap();
let bust_elapsed = bust_start.elapsed();
assert_eq!(third, "server-response-1");
assert!(
bust_elapsed >= Duration::from_millis(100),
"Busted request should still be throttled"
);
let fourth = client.get(&url).send().await.unwrap().text().await.unwrap();
assert_eq!(fourth, "server-response-1");
assert_eq!(request_counter.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_init_client_with_cache_and_throttle() {
let temp_dir = TempDir::new().unwrap();
let cache_path = temp_dir.path().join("client_cache_throttle.bin");
let mock_server = MockServer::start().await;
let response_template = ResponseTemplate::new(200)
.set_body_string("test response")
.insert_header("Cache-Control", "max-age=60");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(response_template.clone())
.expect(1) .mount(&mock_server)
.await;
let cache_policy = CachePolicy {
default_ttl: Duration::from_secs(60), respect_headers: true,
cache_status_override: None,
};
let throttle_policy = ThrottlePolicy {
base_delay_ms: 200, adaptive_jitter_ms: 50, max_concurrent: 1, max_retries: 1, };
let (cache, throttle) = init_cache_with_throttle(&cache_path, cache_policy, throttle_policy);
let client = init_client_with_cache_and_throttle(cache.clone(), throttle.clone());
let url = mock_server.uri();
let start_time_1 = std::time::Instant::now();
let first_response = client.get(&url).send().await.unwrap();
let first_body = first_response.text().await.unwrap();
let elapsed_1 = start_time_1.elapsed();
assert_eq!(first_body, "test response");
assert!(
elapsed_1 >= Duration::from_millis(200),
"First request was too fast!"
);
let start_time_2 = std::time::Instant::now();
let second_response = client.get(&url).send().await.unwrap();
let second_body = second_response.text().await.unwrap();
let elapsed_2 = start_time_2.elapsed();
assert_eq!(second_body, "test response");
assert!(
elapsed_2 < Duration::from_millis(50),
"Second request was not instant despite caching!"
);
tracing::info!(
"Test passed! First request took {:?}, second request took {:?} (should be cached).",
elapsed_1,
elapsed_2
);
}
#[tokio::test]
async fn test_init_throttle_without_store_enforces_base_delay() {
let mock_server = MockServer::start().await;
let response_template = ResponseTemplate::new(200).set_body_string("throttle-only response");
Mock::given(wiremock::matchers::method("GET"))
.respond_with(response_template)
.expect(3)
.mount(&mock_server)
.await;
let throttle_policy = ThrottlePolicy {
base_delay_ms: 120,
adaptive_jitter_ms: 0,
max_concurrent: 1,
max_retries: 0,
};
let throttle = init_throttle(throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle)
.build();
let start = Instant::now();
for i in 0..3 {
let url = format!("{}/throttle-only-{}", mock_server.uri(), i);
let response = client.get(&url).send().await.unwrap();
assert_eq!(response.status(), 200);
assert_eq!(response.text().await.unwrap(), "throttle-only response");
}
let elapsed = start.elapsed();
let min_expected_delay = Duration::from_millis(360);
assert!(
elapsed >= min_expected_delay,
"Throttle-only mode was too fast. Expected at least {:?}, got {:?}",
min_expected_delay,
elapsed
);
}
#[tokio::test]
async fn test_init_throttle_without_store_retries_then_succeeds() {
let mock_server = MockServer::start().await;
let error_template = ResponseTemplate::new(500);
let success_template = ResponseTemplate::new(200).set_body_string("eventual success");
let request_counter = Arc::new(AtomicUsize::new(0));
let counter_clone = Arc::clone(&request_counter);
Mock::given(wiremock::matchers::method("GET"))
.respond_with(move |_: &wiremock::Request| {
let count = counter_clone.fetch_add(1, Ordering::SeqCst);
if count < 2 {
error_template.clone()
} else {
success_template.clone()
}
})
.expect(3)
.mount(&mock_server)
.await;
let throttle_policy = ThrottlePolicy {
base_delay_ms: 100,
adaptive_jitter_ms: 0,
max_concurrent: 1,
max_retries: 2,
};
let throttle = init_throttle(throttle_policy);
let client = ClientBuilder::new(reqwest::Client::new())
.with_arc(throttle)
.build();
let start = Instant::now();
let url = format!("{}/throttle-only-retry", mock_server.uri());
let response = client.get(&url).send().await.unwrap();
let elapsed = start.elapsed();
assert_eq!(response.status(), 200);
assert_eq!(response.text().await.unwrap(), "eventual success");
assert_eq!(request_counter.load(Ordering::SeqCst), 3);
let min_expected = Duration::from_millis(700);
assert!(
elapsed >= min_expected,
"Retry/backoff delay was too short. Expected at least {:?}, got {:?}",
min_expected,
elapsed
);
}