use std::time::Duration;
use http::{HeaderValue, Request, Response, StatusCode, header::HeaderName};
use super::{charge::ChargeMetadata, error::RateLimitError, store::Usage};
const RATE_LIMIT: HeaderName = HeaderName::from_static("ratelimit");
const RATE_LIMIT_POLICY: HeaderName = HeaderName::from_static("ratelimit-policy");
const RETRY_AFTER: HeaderName = HeaderName::from_static("retry-after");
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub enum RateLimitFields {
Draft7,
#[default]
Draft11,
Disabled,
}
#[derive(Debug)]
pub enum ResponseReason {
RateLimited(u64, Usage),
Error(RateLimitError),
}
impl ResponseReason {
pub const fn status_code(&self) -> StatusCode {
match self {
Self::RateLimited(_, _) => StatusCode::TOO_MANY_REQUESTS,
Self::Error(RateLimitError::Key(_, _)) | Self::Error(RateLimitError::Quota(_, _)) => {
StatusCode::INTERNAL_SERVER_ERROR
},
Self::Error(RateLimitError::Store(_, _)) => StatusCode::SERVICE_UNAVAILABLE,
}
}
}
pub trait ResponseFactory<B>: Clone + Send + Sync {
fn build(&self, request: Request<B>, reason: ResponseReason) -> Response<B>;
}
#[derive(Clone, Copy, Debug, Default)]
#[non_exhaustive]
pub struct DefaultResponseFactory;
impl<B> ResponseFactory<B> for DefaultResponseFactory
where
B: Default,
{
fn build(&self, _request: Request<B>, reason: ResponseReason) -> Response<B> {
let mut response = Response::new(B::default());
*response.status_mut() = reason.status_code();
response
}
}
pub(super) enum MiddlewareResponse<B> {
RateLimited(Request<B>, ChargeMetadata),
Error(Request<B>, RateLimitError),
}
impl<B> MiddlewareResponse<B> {
pub(super) fn finalize<F>(self, factory: &F) -> Response<B>
where
F: ResponseFactory<B>,
{
match self {
Self::RateLimited(request, metadata) => {
let reason = ResponseReason::RateLimited(metadata.limit, metadata.usage);
let response = factory.build(request, reason);
append_rate_limited_response_headers(response, metadata)
},
Self::Error(request, error) => factory.build(request, ResponseReason::Error(error)),
}
}
}
pub(super) fn append_inner_response_headers<B>(response: Response<B>, metadata: Option<ChargeMetadata>) -> Response<B> {
match metadata {
Some(metadata) => append_rate_limit_fields(response, &metadata),
None => response,
}
}
fn append_rate_limited_response_headers<B>(response: Response<B>, metadata: ChargeMetadata) -> Response<B> {
let mut response = append_rate_limit_fields(response, &metadata);
append_header(
&mut response,
RETRY_AFTER,
&ceil_seconds(metadata.usage.reset_after).to_string(),
);
response
}
fn append_rate_limit_fields<B>(response: Response<B>, metadata: &ChargeMetadata) -> Response<B> {
let Some((policy, rate_limit)) = format_rate_limit_fields(metadata) else {
return response;
};
let mut response = response;
append_header(&mut response, RATE_LIMIT_POLICY, &policy);
append_header(&mut response, RATE_LIMIT, &rate_limit);
response
}
fn format_rate_limit_fields(metadata: &ChargeMetadata) -> Option<(String, String)> {
if metadata.rate_limit_fields == RateLimitFields::Disabled {
return None;
}
let limit = metadata.limit;
let remaining = metadata.remaining();
let reset_after = ceil_seconds(metadata.usage.reset_after);
let window = ceil_seconds(metadata.window);
Some(match metadata.rate_limit_fields {
RateLimitFields::Draft7 => (
format!("{limit};w={window}"),
format!("limit={limit}, remaining={remaining}, reset={reset_after}"),
),
RateLimitFields::Draft11 => {
let policy_name = &metadata.policy_name;
(
format!(r#""{policy_name}";q={limit};w={window}"#),
format!(r#""{policy_name}";r={remaining};t={reset_after}"#),
)
},
RateLimitFields::Disabled => return None,
})
}
fn append_header<B>(response: &mut Response<B>, name: HeaderName, value: &str) {
if let Ok(value) = HeaderValue::from_str(value) {
response.headers_mut().append(name, value);
}
}
fn ceil_seconds(duration: Duration) -> u64 {
duration
.as_secs()
.saturating_add(u64::from(duration.subsec_nanos() != 0))
.max(1)
}