use std::time::{Duration, SystemTime, UNIX_EPOCH};
use http::header::HeaderValue;
use http::{Response, StatusCode};
use crate::gcra::{ceil_secs, RateLimitInfo, RateLimited};
pub fn inject_headers<B>(response: &mut Response<B>, info: &RateLimitInfo, now: SystemTime) {
let headers = response.headers_mut();
headers.insert("X-RateLimit-Limit", HeaderValue::from(info.limit));
headers.insert("X-RateLimit-Remaining", HeaderValue::from(info.remaining));
headers.insert(
"X-RateLimit-Reset",
reset_header_value(info.reset_after, now),
);
}
pub fn rate_limited_response(
limited: &RateLimited,
tier: &str,
now: SystemTime,
) -> Response<String> {
let retry_after_secs = limited.retry_after_secs();
let escaped_tier = escape_json_string(tier);
let body = format!(
r#"{{"error":"rate limit exceeded","tier":"{}","retry_after":{}}}"#,
escaped_tier, retry_after_secs
);
Response::builder()
.status(StatusCode::TOO_MANY_REQUESTS)
.header("Content-Type", "application/json")
.header("Retry-After", retry_after_secs)
.header("X-RateLimit-Limit", limited.limit)
.header("X-RateLimit-Remaining", 0u32)
.header(
"X-RateLimit-Reset",
reset_header_value(limited.reset_after, now),
)
.body(body)
.unwrap()
}
pub fn cost_exceeds_limit_response(tier: &str, cost: u32, limit: u32) -> Response<String> {
let body = format!(
r#"{{"error":"request cost exceeds tier limit","tier":"{}","cost":{},"limit":{}}}"#,
escape_json_string(tier),
cost,
limit
);
Response::builder()
.status(StatusCode::FORBIDDEN)
.header("Content-Type", "application/json")
.body(body)
.unwrap()
}
pub fn deny_response(status: StatusCode) -> Response<String> {
Response::builder()
.status(status)
.header("Content-Type", "application/json")
.body(format!(r#"{{"error":"{}"}}"#, canonical_reason(status)))
.unwrap()
}
pub fn storage_error_response() -> Response<String> {
Response::builder()
.status(StatusCode::SERVICE_UNAVAILABLE)
.header("Content-Type", "application/json")
.body(r#"{"error":"service unavailable"}"#.to_string())
.unwrap()
}
fn reset_header_value(reset_after: Duration, now: SystemTime) -> HeaderValue {
let secs = match now.checked_add(reset_after) {
Some(reset) => reset.duration_since(UNIX_EPOCH).map_or(0, ceil_secs),
None => u64::MAX,
};
HeaderValue::from(secs)
}
pub fn bad_request_response() -> Response<String> {
Response::builder()
.status(StatusCode::BAD_REQUEST)
.header("Content-Type", "application/json")
.body(r#"{"error":"failed to read request body"}"#.to_string())
.unwrap()
}
fn canonical_reason(status: StatusCode) -> &'static str {
status.canonical_reason().unwrap_or("request denied")
}
fn escape_json_string(s: &str) -> String {
let mut escaped = String::with_capacity(s.len());
for ch in s.chars() {
match ch {
'"' => escaped.push_str("\\\""),
'\\' => escaped.push_str("\\\\"),
c if c.is_control() => {
escaped.push_str(&format!("\\u{:04x}", c as u32));
}
c => escaped.push(c),
}
}
escaped
}
#[cfg(test)]
mod tests {
use super::*;
fn limited(retry_after: Duration) -> RateLimited {
RateLimited {
limit: 1,
retry_after,
reset_after: Duration::from_secs(60),
}
}
#[test]
fn retry_after_is_rounded_up_in_header_and_body() {
let resp =
rate_limited_response(&limited(Duration::from_millis(29_500)), "free", UNIX_EPOCH);
assert_eq!(resp.headers()["retry-after"], "30");
assert!(
resp.body().contains(r#""retry_after":30"#),
"{}",
resp.body()
);
}
#[test]
fn sub_second_retry_after_is_one_in_header_and_body() {
let resp = rate_limited_response(&limited(Duration::from_millis(500)), "free", UNIX_EPOCH);
assert_eq!(resp.headers()["retry-after"], "1");
assert!(
resp.body().contains(r#""retry_after":1"#),
"{}",
resp.body()
);
}
#[test]
fn reset_header_is_unix_time_rounded_up() {
let now = UNIX_EPOCH + Duration::from_secs(1_000);
let info = RateLimitInfo {
limit: 10,
remaining: 9,
reset_after: Duration::from_millis(2_500),
};
let mut resp = Response::new(());
inject_headers(&mut resp, &info, now);
assert_eq!(resp.headers()["x-ratelimit-reset"], "1003");
}
#[test]
fn reset_header_is_zero_when_the_clock_is_before_the_epoch() {
let now = UNIX_EPOCH - Duration::from_secs(10);
let info = RateLimitInfo {
limit: 10,
remaining: 9,
reset_after: Duration::from_secs(1),
};
let mut resp = Response::new(());
inject_headers(&mut resp, &info, now);
assert_eq!(resp.headers()["x-ratelimit-reset"], "0");
}
}