tower_rate_limiter/limiter/
response.rs1use std::time::Duration;
7
8use http::{HeaderValue, Request, Response, StatusCode, header::HeaderName};
9
10use super::{
11 error::RateLimitError,
12 policy::{Policy, ResponseMetadata},
13};
14
15const RATE_LIMIT: HeaderName = HeaderName::from_static("ratelimit");
17const RATE_LIMIT_POLICY: HeaderName = HeaderName::from_static("ratelimit-policy");
19const RETRY_AFTER: HeaderName = HeaderName::from_static("retry-after");
21
22#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
36#[non_exhaustive]
37pub enum RateLimitFields {
38 Draft7,
40 #[default]
42 Draft11,
43 Disabled,
46}
47
48#[derive(Debug)]
50pub enum ResponseReason {
51 RateLimited(Policy),
53 Error(RateLimitError),
55}
56
57impl ResponseReason {
58 pub const fn status_code(&self) -> StatusCode {
60 match self {
61 Self::RateLimited(_) => StatusCode::TOO_MANY_REQUESTS,
62 Self::Error(RateLimitError::Key(_, _)) | Self::Error(RateLimitError::Quota(_, _)) => {
63 StatusCode::INTERNAL_SERVER_ERROR
64 },
65 Self::Error(RateLimitError::Store(_, _)) => StatusCode::SERVICE_UNAVAILABLE,
66 }
67 }
68}
69
70pub trait ResponseFactory<ReqBody, ResBody>: Clone {
72 fn build(&self, request: Request<ReqBody>, reason: ResponseReason) -> Response<ResBody>;
74}
75
76#[derive(Clone, Copy, Debug, Default)]
78#[non_exhaustive]
79pub struct DefaultResponseFactory;
80
81impl<ReqBody, ResBody> ResponseFactory<ReqBody, ResBody> for DefaultResponseFactory
82where
83 ResBody: Default,
84{
85 fn build(&self, _request: Request<ReqBody>, reason: ResponseReason) -> Response<ResBody> {
86 let mut response = Response::new(ResBody::default());
87 *response.status_mut() = reason.status_code();
88 response
89 }
90}
91
92pub(super) enum MiddlewareResponse<ReqBody> {
97 RateLimited(Request<ReqBody>, ResponseMetadata),
98 Error(Request<ReqBody>, RateLimitError),
99}
100
101impl<ReqBody> MiddlewareResponse<ReqBody> {
102 pub(super) fn finalize<ResBody, Factory>(self, factory: &Factory) -> Response<ResBody>
104 where
105 Factory: ResponseFactory<ReqBody, ResBody>,
106 {
107 match self {
108 Self::RateLimited(request, metadata) => {
109 let reason = ResponseReason::RateLimited(metadata.policy.clone());
110 let response = factory.build(request, reason);
111
112 append_rate_limited_response_headers(response, metadata)
113 },
114 Self::Error(request, error) => factory.build(request, ResponseReason::Error(error)),
115 }
116 }
117}
118
119pub(super) fn append_inner_response_headers<B>(
125 response: Response<B>,
126 metadata: Option<ResponseMetadata>,
127) -> Response<B> {
128 match metadata {
129 Some(metadata) => append_rate_limit_fields(response, &metadata),
130 None => response,
131 }
132}
133
134fn append_rate_limited_response_headers<B>(response: Response<B>, metadata: ResponseMetadata) -> Response<B> {
138 let mut response = append_rate_limit_fields(response, &metadata);
139 append_header(
140 &mut response,
141 RETRY_AFTER,
142 &ceil_seconds(metadata.policy.reset_after).to_string(),
143 );
144 response
145}
146
147fn append_rate_limit_fields<B>(response: Response<B>, metadata: &ResponseMetadata) -> Response<B> {
149 let Some((policy, rate_limit)) = format_rate_limit_fields(metadata) else {
150 return response;
151 };
152
153 let mut response = response;
154 append_header(&mut response, RATE_LIMIT_POLICY, &policy);
155 append_header(&mut response, RATE_LIMIT, &rate_limit);
156 response
157}
158
159fn format_rate_limit_fields(metadata: &ResponseMetadata) -> Option<(String, String)> {
161 if metadata.fields == RateLimitFields::Disabled {
162 return None;
163 }
164
165 let limit = metadata.policy.limit;
166 let remaining = metadata.policy.remaining();
167 let reset_after = ceil_seconds(metadata.policy.reset_after);
168 let window = ceil_seconds(metadata.policy.window);
169
170 Some(match metadata.fields {
171 RateLimitFields::Draft7 => (
172 format!("{limit};w={window}"),
173 format!("limit={limit}, remaining={remaining}, reset={reset_after}"),
174 ),
175 RateLimitFields::Draft11 => {
176 let policy_name = &metadata.policy.name;
177
178 (
179 format!(r#""{policy_name}";q={limit};w={window}"#),
180 format!(r#""{policy_name}";r={remaining};t={reset_after}"#),
181 )
182 },
183 RateLimitFields::Disabled => return None,
184 })
185}
186
187fn append_header<B>(response: &mut Response<B>, name: HeaderName, value: &str) {
189 if let Ok(value) = HeaderValue::from_str(value) {
190 response.headers_mut().append(name, value);
191 }
192}
193
194fn ceil_seconds(duration: Duration) -> u64 {
196 duration
197 .as_secs()
198 .saturating_add(u64::from(duration.subsec_nanos() != 0))
199 .max(1)
200}