#![allow(dead_code)]
use std::future::Future;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use tokio::task::JoinHandle;
use vantage_api_pool::resilient::{
AuthRefresher, BreakerMode, CallPolicy, ErrorKind, RetryMode, RetryPolicy, TransportEvent,
TransportObserver,
};
use vantage_api_pool::ResilientClient;
use wiremock::matchers::method;
use wiremock::{Mock, MockServer, Request, Respond, ResponseTemplate};
pub fn fast_retry(max_retries: usize) -> RetryPolicy {
RetryPolicy {
max_retries,
base_backoff: Duration::from_millis(1),
max_backoff: Duration::from_millis(5),
}
}
pub fn essential_fast() -> CallPolicy {
CallPolicy {
retry: RetryMode::UntilCancelled {
base: Duration::from_millis(1),
max: Duration::from_millis(4),
},
breaker: BreakerMode::WaitForProbe,
}
}
pub async fn within<T>(fut: impl Future<Output = T>) -> T {
within_secs(5, fut).await
}
pub async fn within_secs<T>(secs: u64, fut: impl Future<Output = T>) -> T {
tokio::time::timeout(Duration::from_secs(secs), fut)
.await
.unwrap_or_else(|_| panic!("the call did not finish within {secs}s"))
}
pub fn spawn_essential(client: &ResilientClient, url: &str) -> JoinHandle<()> {
let client = client.clone();
let url = url.to_string();
tokio::spawn(async move {
let _ = within(client.execute_with(&essential_fast(), |h| h.get(&url))).await;
})
}
pub async fn cancel(task: JoinHandle<()>) {
task.abort();
let _ = task.await;
}
pub async fn background_error(client: &ResilientClient, url: &str) -> ErrorKind {
client
.execute_with(&CallPolicy::background(), |h| h.get(url))
.await
.expect_err("the call was expected to fail")
.kind
}
pub async fn background_ok(client: &ResilientClient, url: &str) -> u16 {
client
.execute_with(&CallPolicy::background(), |h| h.get(url))
.await
.expect("the call was expected to succeed")
.status()
.as_u16()
}
pub async fn open_breaker(client: &ResilientClient, url: &str, failures: usize) {
for _ in 0..failures {
let _ = client
.execute_with(&CallPolicy::background(), |h| h.get(url))
.await;
}
}
pub async fn mount(server: &MockServer, responder: impl Respond + 'static) {
Mock::given(method("GET"))
.respond_with(responder)
.mount(server)
.await;
}
pub async fn mount_status(server: &MockServer, status: u16) {
mount(server, ResponseTemplate::new(status)).await;
}
pub async fn mount_script(server: &MockServer, statuses: impl Into<Vec<u16>>) {
mount(server, Script::new(statuses)).await;
}
pub async fn mount_delayed(server: &MockServer, status: u16, delay: Duration) {
mount(server, ResponseTemplate::new(status).set_delay(delay)).await;
}
pub async fn requests(server: &MockServer) -> usize {
server.received_requests().await.unwrap().len()
}
pub async fn wait_for_requests(server: &MockServer, n: usize, limit: Duration) -> usize {
let deadline = Instant::now() + limit;
loop {
let seen = requests(server).await;
if seen >= n {
return seen;
}
assert!(
Instant::now() < deadline,
"the server saw {seen} of {n} requests within {limit:?}"
);
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
pub async fn wait_until(limit: Duration, what: &str, cond: impl Fn() -> bool) {
let deadline = Instant::now() + limit;
while !cond() {
assert!(
Instant::now() < deadline,
"{what} did not happen in {limit:?}"
);
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
#[derive(Default)]
pub struct HalfFailing(AtomicUsize);
impl Respond for HalfFailing {
fn respond(&self, _request: &Request) -> ResponseTemplate {
let n = self.0.fetch_add(1, Ordering::SeqCst);
ResponseTemplate::new(if n.is_multiple_of(2) { 503 } else { 200 })
}
}
pub struct Script {
statuses: Vec<u16>,
next: AtomicUsize,
}
impl Script {
pub fn new(statuses: impl Into<Vec<u16>>) -> Self {
let statuses = statuses.into();
assert!(!statuses.is_empty(), "a script needs at least one status");
Self {
statuses,
next: AtomicUsize::new(0),
}
}
}
impl Respond for Script {
fn respond(&self, _request: &Request) -> ResponseTemplate {
let i = self
.next
.fetch_add(1, Ordering::SeqCst)
.min(self.statuses.len() - 1);
ResponseTemplate::new(self.statuses[i])
}
}
#[derive(Clone, Default)]
pub struct Calls(Arc<AtomicUsize>);
impl Calls {
pub fn count(&self) -> usize {
self.0.load(Ordering::SeqCst)
}
}
pub fn scripted_refresher(
answer: impl Fn(usize) -> anyhow::Result<String> + Send + Sync + 'static,
) -> (AuthRefresher, Calls) {
let calls = Calls::default();
let counter = calls.clone();
let refresher: AuthRefresher = Arc::new(move || {
let token = answer(counter.0.fetch_add(1, Ordering::SeqCst));
Box::pin(async move { token })
});
(refresher, calls)
}
pub fn counting_refresher() -> (AuthRefresher, Calls) {
scripted_refresher(|n| Ok(format!("token-{n}")))
}
#[derive(Default)]
pub struct Recorder {
events: Mutex<Vec<(String, String)>>,
failed_ms: Mutex<Vec<u64>>,
succeeded_ms: Mutex<Vec<u64>>,
}
impl Recorder {
pub fn tags(&self) -> Vec<String> {
self.events
.lock()
.unwrap()
.iter()
.map(|(_, tag)| tag.clone())
.collect()
}
pub fn key(&self) -> Option<String> {
let events = self.events.lock().unwrap();
let mut keys = events.iter().map(|(k, _)| k);
let first = keys.next()?.clone();
assert!(keys.all(|k| *k == first), "events used more than one key");
Some(first)
}
pub fn failed_ms(&self) -> Vec<u64> {
self.failed_ms.lock().unwrap().clone()
}
pub fn succeeded_ms(&self) -> Vec<u64> {
self.succeeded_ms.lock().unwrap().clone()
}
}
impl TransportObserver for Recorder {
fn on_event(&self, key: &str, event: TransportEvent) {
let tag = match event {
TransportEvent::Started => "started".to_string(),
TransportEvent::Succeeded { status, ms, .. } => {
self.succeeded_ms.lock().unwrap().push(ms);
format!("ok:{status}")
}
TransportEvent::Failed { error, ms } => {
self.failed_ms.lock().unwrap().push(ms);
format!("failed:{}", error.kind_name())
}
TransportEvent::Cancelled => "cancelled".to_string(),
TransportEvent::RetryScheduled { attempt, .. } => format!("retry:{attempt}"),
TransportEvent::BreakerOpened { .. } => "opened".to_string(),
TransportEvent::BreakerClosed => "closed".to_string(),
TransportEvent::RowsPulled { n } => format!("rows:{n}"),
TransportEvent::WritePushed => "write".to_string(),
};
self.events.lock().unwrap().push((key.to_string(), tag));
}
}