use http::Extensions;
use keel_core_api::{AttemptResult, ENVELOPE_VERSION, ErrorClass, Request as CoreRequest};
use reqwest::{Request, Response};
use reqwest_middleware::{Error as MwError, Middleware, Next, Result as MwResult};
use std::sync::{Arc, Mutex};
const IDEMPOTENT_METHODS: [&str; 6] = ["GET", "HEAD", "OPTIONS", "PUT", "DELETE", "TRACE"];
const DEFAULT_IDEMPOTENCY_HEADERS: [&str; 2] = ["idempotency-key", "x-idempotency-key"];
#[derive(Debug, Clone)]
pub struct KeelMiddleware {
client: reqwest::Client,
}
impl KeelMiddleware {
#[must_use]
pub fn new(client: reqwest::Client) -> Self {
Self { client }
}
}
#[derive(Clone)]
struct LiveState {
pending: Arc<Mutex<Option<Request>>>,
ok: Arc<Mutex<Option<Response>>>,
transient: Arc<Mutex<Option<Response>>>,
err: Arc<Mutex<Option<MwError>>>,
}
impl LiveState {
fn new(req: Request) -> Self {
Self {
pending: Arc::new(Mutex::new(Some(req))),
ok: Arc::new(Mutex::new(None)),
transient: Arc::new(Mutex::new(None)),
err: Arc::new(Mutex::new(None)),
}
}
}
async fn run_attempt(client: reqwest::Client, live: LiveState) -> AttemptResult {
let attempt_req = match take_attempt_request(&live.pending) {
Ok(req) => req,
Err(err) => return err,
};
match client.execute(attempt_req).await {
Ok(resp) => {
let status = resp.status().as_u16();
if is_transient_status(status) {
let retry_after_ms = parse_retry_after_ms(resp.headers());
let message = format!("http {status}");
*live.transient.lock().expect("keel: mutex poisoned") = Some(resp);
AttemptResult::Error {
class: ErrorClass::Http,
http_status: Some(status),
retry_after_ms,
message,
original: None,
}
} else {
*live.ok.lock().expect("keel: mutex poisoned") = Some(resp);
AttemptResult::Ok {
payload: serde_json::Value::Null,
}
}
}
Err(err) => {
let class = classify(&err);
let message = err.to_string();
*live.err.lock().expect("keel: mutex poisoned") = Some(err.into());
AttemptResult::Error {
class,
http_status: None,
retry_after_ms: None,
message,
original: None,
}
}
}
}
fn take_attempt_request(pending: &Mutex<Option<Request>>) -> Result<Request, AttemptResult> {
let mut guard = pending.lock().expect("keel: pending mutex poisoned");
if let Some(cloned) = guard.as_ref().and_then(Request::try_clone) {
return Ok(cloned);
}
guard.take().ok_or_else(|| AttemptResult::Error {
class: ErrorClass::Other,
http_status: None,
retry_after_ms: None,
message: "keel: request body cannot be cloned for a retry attempt".to_owned(),
original: None,
})
}
fn deliver(outcome: &keel_core_api::Outcome, live: &LiveState) -> MwResult<Response> {
if outcome.result == "ok" {
let resp = live
.ok
.lock()
.expect("keel: mutex poisoned")
.take()
.expect("keel: ok outcome without a live response");
return Ok(resp);
}
if let Some(resp) = live.transient.lock().expect("keel: mutex poisoned").take() {
return Ok(resp);
}
if let Some(err) = live.err.lock().expect("keel: mutex poisoned").take() {
return Err(err);
}
let outcome_error = outcome
.error
.clone()
.expect("engine reported an error outcome without an OutcomeError");
Err(MwError::middleware(
crate::Error::<std::convert::Infallible>::Keel(outcome_error),
))
}
#[async_trait::async_trait]
impl Middleware for KeelMiddleware {
async fn handle(
&self,
req: Request,
_extensions: &mut Extensions,
_next: Next<'_>,
) -> MwResult<Response> {
let method = req.method().clone();
let host = req.url().host_str().unwrap_or("unknown").to_owned();
let op = format!("{method} {host}{}", req.url().path());
let idempotent = is_idempotent(method.as_str(), req.headers());
let core_req = CoreRequest {
v: ENVELOPE_VERSION,
target: host,
op,
idempotent,
args_hash: None,
};
let client = self.client.clone();
let live = LiveState::new(req);
let outcome = crate::engine()
.execute(&core_req, {
let live = live.clone();
move |_attempt: u32| run_attempt(client.clone(), live.clone())
})
.await;
deliver(&outcome, &live)
}
}
fn is_idempotent(method: &str, headers: &reqwest::header::HeaderMap) -> bool {
if IDEMPOTENT_METHODS.contains(&method) {
return true;
}
headers
.keys()
.any(|name| DEFAULT_IDEMPOTENCY_HEADERS.contains(&name.as_str()))
}
fn is_transient_status(status: u16) -> bool {
status == 429 || (500..=599).contains(&status)
}
fn parse_retry_after_ms(headers: &reqwest::header::HeaderMap) -> Option<u64> {
let value = headers.get(reqwest::header::RETRY_AFTER)?.to_str().ok()?;
let secs: u64 = value.trim().parse().ok()?;
Some(secs.saturating_mul(1000))
}
fn classify(err: &reqwest::Error) -> ErrorClass {
if err.is_timeout() {
ErrorClass::Timeout
} else if err.is_connect() {
ErrorClass::Conn
} else {
ErrorClass::Other
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn idempotent_methods() {
for m in ["GET", "HEAD", "OPTIONS", "PUT", "DELETE", "TRACE"] {
assert!(is_idempotent(m, &reqwest::header::HeaderMap::new()), "{m}");
}
assert!(!is_idempotent("POST", &reqwest::header::HeaderMap::new()));
assert!(!is_idempotent("PATCH", &reqwest::header::HeaderMap::new()));
}
#[test]
fn post_with_idempotency_key_header_is_idempotent() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert("idempotency-key", "abc".parse().unwrap());
assert!(is_idempotent("POST", &headers));
}
#[test]
fn transient_status() {
assert!(is_transient_status(429));
assert!(is_transient_status(500));
assert!(is_transient_status(503));
assert!(!is_transient_status(404));
assert!(!is_transient_status(200));
assert!(!is_transient_status(301));
}
#[test]
fn retry_after_seconds_form() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(reqwest::header::RETRY_AFTER, "2".parse().unwrap());
assert_eq!(parse_retry_after_ms(&headers), Some(2000));
}
#[test]
fn retry_after_http_date_form_is_unparsed() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::RETRY_AFTER,
"Wed, 21 Oct 2026 07:28:00 GMT".parse().unwrap(),
);
assert_eq!(parse_retry_after_ms(&headers), None);
}
}