use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::SystemTime;
use bytes::Bytes;
use http::{Request, Response, StatusCode};
use http_body::Body;
use http_body_util::{BodyExt, Full, LengthLimitError, Limited};
use tower_layer::Layer;
use tower_service::Service;
use crate::check::{self, CheckOutcome};
use crate::layer::Settings;
use crate::response;
use crate::storage::StorageKey;
use crate::tier::RateTier;
#[derive(Clone, Debug)]
pub struct BufferedTierLimitLayer {
pub(crate) rate_tier: Arc<RateTier>,
pub(crate) settings: Settings,
pub(crate) max_body_size: usize,
}
impl BufferedTierLimitLayer {
pub fn max_body_size(mut self, size: usize) -> Self {
self.max_body_size = size;
self
}
}
impl<S> Layer<S> for BufferedTierLimitLayer {
type Service = BufferedTierLimitService<S>;
fn layer(&self, inner: S) -> Self::Service {
BufferedTierLimitService {
inner,
rate_tier: self.rate_tier.clone(),
settings: Arc::new(self.settings.clone()),
max_body_size: self.max_body_size,
}
}
}
#[derive(Clone, Debug)]
pub struct BufferedTierLimitService<S> {
inner: S,
rate_tier: Arc<RateTier>,
settings: Arc<Settings>,
max_body_size: usize,
}
impl<S, B, ResBody> Service<Request<B>> for BufferedTierLimitService<S>
where
S: Service<Request<Full<Bytes>>, Response = Response<ResBody>> + Clone + Send + 'static,
S::Future: Send,
S::Error: Send,
B: Body + Send + 'static,
B::Data: Send,
B::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
ResBody: From<String> + Send,
{
type Response = Response<ResBody>;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request<B>) -> Self::Future {
let rate_tier = self.rate_tier.clone();
let settings = self.settings.clone();
let max_body_size = self.max_body_size;
let mut inner = self.inner.clone();
std::mem::swap(&mut self.inner, &mut inner);
Box::pin(async move {
let (parts, body) = req.into_parts();
if body.size_hint().lower() > max_body_size as u64 {
return Ok(payload_too_large_response().map(Into::into));
}
let body_bytes = match Limited::new(body, max_body_size).collect().await {
Ok(collected) => collected.to_bytes(),
Err(err) if err.is::<LengthLimitError>() => {
return Ok(payload_too_large_response().map(Into::into));
}
Err(_) => {
return Ok(response::bad_request_response().map(Into::into));
}
};
let identity = settings
.identifier
.identify_with_body(&parts.headers, &body_bytes)
.await;
let (user_id, tier_name) = match check::resolve_identity(identity, &rate_tier) {
Ok(pair) => pair,
Err(CheckOutcome::PassThrough) => {
let req = Request::from_parts(parts, Full::new(body_bytes));
return inner.call(req).await;
}
Err(CheckOutcome::Deny(resp)) => return Ok(resp.map(Into::into)),
Err(CheckOutcome::Allow(_)) => unreachable!(),
};
let resolved = check::resolve_quota(&rate_tier, &user_id, tier_name, &settings);
let (tier_name, quota) = match resolved {
Ok(resolved) => resolved,
Err(CheckOutcome::PassThrough) => {
let req = Request::from_parts(parts, Full::new(body_bytes));
return inner.call(req).await;
}
Err(CheckOutcome::Deny(resp)) => return Ok(resp.map(Into::into)),
Err(CheckOutcome::Allow(_)) => unreachable!(),
};
let cost = check::request_cost(&parts, &settings);
if let Some(resp) =
check::reject_cost_over_limit(cost, quota, &user_id, &tier_name, &settings)
{
return Ok(resp.map(Into::into));
}
let now = rate_tier.clock().now();
let key = StorageKey::new(&user_id, &tier_name);
let result = rate_tier
.storage()
.check_and_update(key, quota, cost, now)
.await;
let checked_at = SystemTime::now();
match check::process_result(result, &user_id, &tier_name, &settings, checked_at) {
CheckOutcome::Allow(info) => {
let req = Request::from_parts(parts, Full::new(body_bytes));
let mut resp = inner.call(req).await?;
response::inject_headers(&mut resp, &info, checked_at);
Ok(resp)
}
CheckOutcome::Deny(resp) => Ok(resp.map(Into::into)),
CheckOutcome::PassThrough => {
let req = Request::from_parts(parts, Full::new(body_bytes));
inner.call(req).await
}
}
})
}
}
fn payload_too_large_response() -> Response<String> {
Response::builder()
.status(StatusCode::PAYLOAD_TOO_LARGE)
.header("Content-Type", "application/json")
.body(r#"{"error":"payload too large"}"#.to_string())
.unwrap()
}