use std::sync::Arc;
use axum::body::Body;
use axum::http::{HeaderValue, Response, StatusCode, header};
use tower_governor::governor::GovernorConfigBuilder;
use tower_governor::{GovernorError, GovernorLayer};
use utoipa_axum::router::OpenApiRouter;
use vta_config::ServerConfig;
pub const RATE_LIMIT_SOURCE_HEADER: &str = vta_sdk::rate_limit::SOURCE_HEADER;
pub const RATE_LIMIT_SOURCE_VTA: &str = "vta";
pub const RATE_LIMIT_SCOPE_HEADER: &str = "x-rate-limit-scope";
const LEGACY_RATE_LIMIT_AFTER_HEADER: &str = vta_sdk::rate_limit::LEGACY_RETRY_AFTER_HEADER;
pub(crate) const AUTH_INTERVAL_SECS: u64 = 5;
pub(crate) const AUTH_BURST: u32 = 10;
pub(crate) const DID_LOG_INTERVAL_SECS: u64 = 1;
pub(crate) const DID_LOG_BURST: u32 = 60;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Limiter {
Auth,
DidLog,
BackupBlob,
}
impl Limiter {
pub const fn name(self) -> &'static str {
match self {
Limiter::Auth => "auth",
Limiter::DidLog => "did-log",
Limiter::BackupBlob => "backup-blob",
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Quota {
interval_secs: u64,
burst: u32,
}
impl Quota {
pub fn new(interval_secs: u64, burst: u32) -> Self {
Self {
interval_secs: interval_secs.max(1),
burst: burst.max(1),
}
}
pub fn interval_secs(self) -> u64 {
self.interval_secs
}
pub fn burst(self) -> u32 {
self.burst
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RateLimits {
auth: Quota,
did_log: Quota,
}
impl RateLimits {
pub fn new(auth: Quota, did_log: Quota) -> Self {
Self { auth, did_log }
}
pub fn from_server_config(server: &ServerConfig) -> Self {
Self::new(
Quota::new(server.rate_limit_interval_secs, server.rate_limit_burst),
Quota::new(
server.did_log_rate_limit_interval_secs,
server.did_log_rate_limit_burst,
),
)
}
pub fn quota(self, limiter: Limiter) -> Quota {
match limiter {
Limiter::Auth | Limiter::BackupBlob => self.auth,
Limiter::DidLog => self.did_log,
}
}
}
impl Default for RateLimits {
fn default() -> Self {
Self::new(
Quota::new(AUTH_INTERVAL_SECS, AUTH_BURST),
Quota::new(DID_LOG_INTERVAL_SECS, DID_LOG_BURST),
)
}
}
pub(super) fn apply<S>(
router: OpenApiRouter<S>,
limiter: Limiter,
trust_xff: bool,
limits: RateLimits,
) -> OpenApiRouter<S>
where
S: Clone + Send + Sync + 'static,
{
let quota = limits.quota(limiter);
let on_error = move |err: GovernorError| governor_error_response(limiter, err);
if trust_xff {
let cfg = Arc::new(
GovernorConfigBuilder::default()
.per_second(quota.interval_secs())
.burst_size(quota.burst())
.key_extractor(tower_governor::key_extractor::SmartIpKeyExtractor)
.finish()
.expect("Quota clamps interval and burst to >= 1"),
);
router.layer(GovernorLayer::new(cfg).error_handler(on_error))
} else {
let cfg = Arc::new(
GovernorConfigBuilder::default()
.per_second(quota.interval_secs())
.burst_size(quota.burst())
.key_extractor(tower_governor::key_extractor::PeerIpKeyExtractor)
.finish()
.expect("Quota clamps interval and burst to >= 1"),
);
router.layer(GovernorLayer::new(cfg).error_handler(on_error))
}
}
fn governor_error_response(limiter: Limiter, err: GovernorError) -> Response<Body> {
match err {
GovernorError::TooManyRequests { wait_time, .. } => {
too_many_requests(limiter, wait_time.saturating_add(1))
}
other => other.into_response().map(Body::from),
}
}
pub(crate) fn too_many_requests(limiter: Limiter, retry_after_secs: u64) -> Response<Body> {
let retry_after = retry_after_secs.max(1);
let body = serde_json::json!({
"error": "rate_limited",
"limiter": limiter.name(),
"message": format!(
"Too Many Requests: rejected by the VTA's `{}` rate limiter. Retry after {retry_after} s.",
limiter.name()
),
"retryAfterSecs": retry_after,
})
.to_string();
Response::builder()
.status(StatusCode::TOO_MANY_REQUESTS)
.header(header::CONTENT_TYPE, "application/json")
.header(header::RETRY_AFTER, HeaderValue::from(retry_after))
.header(
LEGACY_RATE_LIMIT_AFTER_HEADER,
HeaderValue::from(retry_after),
)
.header(RATE_LIMIT_SOURCE_HEADER, RATE_LIMIT_SOURCE_VTA)
.header(RATE_LIMIT_SCOPE_HEADER, limiter.name())
.body(Body::from(body))
.expect("static header names and numeric values always build")
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::to_bytes;
use axum::http::Request;
use axum::routing::get;
use tower::ServiceExt;
async fn ok() -> &'static str {
"ok"
}
fn two_branch_router(limits: RateLimits) -> axum::Router {
let auth = apply(
OpenApiRouter::<()>::new().route("/auth", get(ok)),
Limiter::Auth,
true,
limits,
);
let did_log = apply(
OpenApiRouter::<()>::new().route("/did.jsonl", get(ok)),
Limiter::DidLog,
true,
limits,
);
let (router, _) = OpenApiRouter::<()>::new()
.merge(auth)
.merge(did_log)
.split_for_parts();
router
}
async fn get_status(app: &axum::Router, uri: &str, ip: &str) -> Response<Body> {
let req = Request::builder()
.uri(uri)
.header("x-forwarded-for", ip)
.body(Body::empty())
.unwrap();
app.clone().oneshot(req).await.unwrap()
}
fn tight() -> RateLimits {
RateLimits::new(Quota::new(3600, 2), Quota::new(3600, 3))
}
#[tokio::test]
async fn did_log_burst_does_not_spend_auth_budget() {
let app = two_branch_router(tight());
for _ in 0..3 {
let r = get_status(&app, "/did.jsonl", "198.51.100.1").await;
assert_eq!(r.status(), StatusCode::OK);
}
let r = get_status(&app, "/did.jsonl", "198.51.100.1").await;
assert_eq!(r.status(), StatusCode::TOO_MANY_REQUESTS);
for _ in 0..2 {
let r = get_status(&app, "/auth", "198.51.100.1").await;
assert_eq!(
r.status(),
StatusCode::OK,
"exhausting did-log must not spend the auth bucket"
);
}
}
#[tokio::test]
async fn auth_burst_does_not_spend_did_log_budget() {
let app = two_branch_router(tight());
for _ in 0..2 {
assert_eq!(
get_status(&app, "/auth", "198.51.100.2").await.status(),
StatusCode::OK
);
}
assert_eq!(
get_status(&app, "/auth", "198.51.100.2").await.status(),
StatusCode::TOO_MANY_REQUESTS
);
for _ in 0..3 {
assert_eq!(
get_status(&app, "/did.jsonl", "198.51.100.2")
.await
.status(),
StatusCode::OK,
"exhausting auth must not spend the did-log bucket"
);
}
}
#[tokio::test]
async fn limits_are_per_ip() {
let app = two_branch_router(tight());
for _ in 0..2 {
get_status(&app, "/auth", "198.51.100.3").await;
}
assert_eq!(
get_status(&app, "/auth", "198.51.100.3").await.status(),
StatusCode::TOO_MANY_REQUESTS
);
assert_eq!(
get_status(&app, "/auth", "198.51.100.4").await.status(),
StatusCode::OK
);
}
#[tokio::test]
async fn rejection_carries_the_vta_429_contract() {
let app = two_branch_router(tight());
for (uri, scope, n) in [("/auth", "auth", 2), ("/did.jsonl", "did-log", 3)] {
for _ in 0..n {
get_status(&app, uri, "198.51.100.5").await;
}
let r = get_status(&app, uri, "198.51.100.5").await;
assert_eq!(r.status(), StatusCode::TOO_MANY_REQUESTS);
let h = r.headers();
assert_eq!(h[RATE_LIMIT_SOURCE_HEADER], "vta");
assert_eq!(h[RATE_LIMIT_SCOPE_HEADER], scope);
let retry: u64 = h[header::RETRY_AFTER].to_str().unwrap().parse().unwrap();
assert!(retry >= 1, "retry-after must be at least 1 s, got {retry}");
assert_eq!(h[LEGACY_RATE_LIMIT_AFTER_HEADER], h[header::RETRY_AFTER]);
assert_eq!(h[header::CONTENT_TYPE], "application/json");
let body = to_bytes(r.into_body(), usize::MAX).await.unwrap();
let body: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(
body,
serde_json::json!({
"error": "rate_limited",
"limiter": scope,
"message": format!(
"Too Many Requests: rejected by the VTA's `{scope}` rate limiter. \
Retry after {retry} s."
),
"retryAfterSecs": retry,
})
);
}
}
#[test]
fn source_label_is_what_the_sdk_reads_as_vta() {
use vta_sdk::rate_limit::RateLimitSource;
assert_eq!(
RateLimitSource::from_source_header(Some(RATE_LIMIT_SOURCE_VTA)),
RateLimitSource::Vta
);
assert_eq!(RATE_LIMIT_SOURCE_HEADER, "x-rate-limit-source");
}
#[test]
fn scope_names_are_stable() {
assert_eq!(Limiter::Auth.name(), "auth");
assert_eq!(Limiter::DidLog.name(), "did-log");
assert_eq!(Limiter::BackupBlob.name(), "backup-blob");
}
#[test]
fn zero_quota_is_clamped() {
let q = Quota::new(0, 0);
assert_eq!((q.interval_secs(), q.burst()), (1, 1));
let server = ServerConfig {
rate_limit_interval_secs: 0,
rate_limit_burst: 0,
did_log_rate_limit_interval_secs: 0,
did_log_rate_limit_burst: 0,
..Default::default()
};
let limits = RateLimits::from_server_config(&server);
assert_eq!(limits.quota(Limiter::DidLog), Quota::new(1, 1));
assert_eq!(limits.quota(Limiter::Auth), Quota::new(1, 1));
let _ = two_branch_router(limits);
}
#[test]
fn defaults_match_server_config_defaults() {
assert_eq!(
RateLimits::default(),
RateLimits::from_server_config(&ServerConfig::default()),
"the no-config defaults must match vta-config's [server] defaults"
);
}
#[test]
fn backup_blob_uses_the_auth_quota() {
let limits = tight();
assert_eq!(
limits.quota(Limiter::BackupBlob),
limits.quota(Limiter::Auth)
);
}
}