pub mod gcra;
#[cfg(feature = "redis")]
pub mod global;
pub mod matcher;
pub mod rate;
pub mod resolve;
pub mod store;
pub mod window;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use axum::extract::Request;
use axum::http::{HeaderMap, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use axum::Json;
use crate::config::ShieldConfig;
use gcra::Verdict;
use matcher::{CompiledProfile, CompiledRule, KeySource, Phase};
use store::GcraStore;
pub struct Shield {
rules: Vec<CompiledRule>,
profiles: HashMap<String, CompiledProfile>,
default_profile: Option<CompiledProfile>,
jwt_limits: Option<resolve::JwtLimits>,
limit_service: Option<Arc<resolve::LimitService>>,
#[cfg(feature = "redis")]
global: Option<Arc<global::GlobalCounters>>,
store: GcraStore,
trusted_proxies: Vec<ipnet::IpNet>,
}
impl Shield {
pub fn build(config: &ShieldConfig) -> Result<Option<Arc<Self>>, String> {
if !config.enabled {
return Ok(None);
}
if config.rules.is_empty() {
return Err(
"shield.enabled is true but no rules are configured (note the schema: \
profiles + rules + sync, not the older endpoint_classes/identifier_endpoints)"
.to_string(),
);
}
let profiles = matcher::compile_profiles(&config.profiles)?;
let rules = matcher::compile_rules(&config.rules, &profiles)?;
let default_profile =
match &config.default_profile {
Some(name) => Some(*profiles.get(name).ok_or_else(|| {
format!("default_profile references unknown profile {name:?}")
})?),
None => None,
};
let trusted_proxies = config
.trusted_proxies
.iter()
.map(|s| parse_cidr(s))
.collect::<Result<Vec<_>, _>>()?;
let jwt_limits = config
.jwt_limits
.as_ref()
.map(resolve::JwtLimits::from_config);
let limit_service = match &config.limit_service {
Some(cfg) => Some(resolve::LimitService::build(cfg, profiles.clone())?),
None => None,
};
#[cfg(feature = "redis")]
let global = match &config.sync {
Some(sync) => {
let g = global::GlobalCounters::build(
&sync.redis_url,
Duration::from_millis(sync.interval_ms),
)?;
g.spawn().then_some(g)
}
None => None,
};
#[cfg(not(feature = "redis"))]
if config.sync.is_some() {
tracing::warn!(
"shield.sync is set but the `redis` feature is not compiled in; \
staying local-only (per-instance limits)"
);
}
Ok(Some(Arc::new(Self {
rules,
profiles,
default_profile,
jwt_limits,
limit_service,
#[cfg(feature = "redis")]
global,
store: GcraStore::new(),
trusted_proxies,
})))
}
fn match_rule(&self, path: &str, phase: Phase) -> Option<&CompiledRule> {
self.rules
.iter()
.find(|r| r.phase == phase && r.matcher.is_match(path))
}
fn resolve_limit(
&self,
rule: &CompiledRule,
claims: Option<&serde_json::Value>,
identity: &str,
) -> Option<CompiledProfile> {
if let (Some(jwt), Some(claims)) = (&self.jwt_limits, claims) {
if let Some(profile) = jwt.resolve(claims, &self.profiles) {
return Some(profile);
}
}
if let Some(service) = &self.limit_service {
if let Some(profile) = service.resolve(identity) {
return Some(profile);
}
}
self.static_profile(rule)
}
fn static_profile(&self, rule: &CompiledRule) -> Option<CompiledProfile> {
rule.profile
.as_ref()
.and_then(|name| self.profiles.get(name))
.or(self.default_profile.as_ref())
.copied()
}
}
pub async fn pre_auth_middleware(
axum::extract::State(shield): axum::extract::State<Arc<Shield>>,
request: Request,
next: Next,
) -> Response {
enforce(&shield, Phase::PreAuth, request, next).await
}
pub async fn post_auth_middleware(
axum::extract::State(shield): axum::extract::State<Arc<Shield>>,
request: Request,
next: Next,
) -> Response {
enforce(&shield, Phase::PostAuth, request, next).await
}
async fn enforce(shield: &Shield, phase: Phase, request: Request, next: Next) -> Response {
let path = request.uri().path();
let Some(rule) = shield.match_rule(path, phase) else {
return next.run(request).await;
};
let peer = request
.extensions()
.get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
.map(|ci| ci.0.ip());
let client = client_ip(peer, request.headers(), &shield.trusted_proxies);
let claims = request
.extensions()
.get::<crate::auth::ValidatedClaims>()
.map(|c| c.0.as_ref());
let key = rule_key(
&rule.fingerprint,
&rule.key,
&client,
request.headers(),
claims,
);
let Some(profile) = shield.resolve_limit(rule, claims, &key.identity) else {
return next.run(request).await;
};
#[cfg(feature = "redis")]
let fleet_remaining = shield
.global
.as_ref()
.map(|g| g.fleet_remaining(&key.store, profile.limit, profile.window));
#[cfg(feature = "redis")]
if fleet_remaining == Some(0) {
return global_reject(profile.limit, profile.window);
}
let verdict = shield.store.check(&key.store, &profile.gcra);
if !verdict.allowed {
return too_many_requests(profile.limit, verdict.remaining, &verdict);
}
#[cfg(not(feature = "redis"))]
let (reported, header_verdict) = reconciled_headers(verdict, None, profile.window);
#[cfg(feature = "redis")]
let (reported, header_verdict) = {
if fleet_remaining.is_some() {
if let Some(global) = &shield.global {
global.record(&key.store, profile.window);
}
}
reconciled_headers(verdict, fleet_remaining, profile.window)
};
let mut response = next.run(request).await;
maybe_tighten_rate_headers(
response.headers_mut(),
profile.limit,
reported,
&header_verdict,
);
response
}
fn reconciled_headers(
verdict: Verdict,
fleet_remaining: Option<u64>,
window: Duration,
) -> (u64, Verdict) {
let mut reported = verdict.remaining;
let mut hv = verdict;
if let Some(fr) = fleet_remaining {
let fleet_r = fr.saturating_sub(1);
if fleet_r <= reported {
reported = fleet_r;
let backoff = fleet_backoff(window);
hv.reset_after = hv.reset_after.max(backoff);
hv.retry_after = hv.retry_after.max(backoff);
}
}
(reported, hv)
}
fn fleet_backoff(window: Duration) -> Duration {
(window / 10).max(Duration::from_secs(1))
}
fn maybe_tighten_rate_headers(
headers: &mut HeaderMap,
limit: u64,
remaining: u64,
verdict: &Verdict,
) {
let header_u64 = |name: &str| {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u64>().ok())
};
let keep_inner = match header_u64("ratelimit-remaining") {
Some(inner) if inner < remaining => true,
Some(inner) if inner == remaining => {
header_u64("ratelimit-reset").unwrap_or(0) >= secs_ceil(verdict.reset_after)
}
_ => false,
};
if !keep_inner {
attach_rate_headers(headers, limit, remaining, verdict);
if let Some(current) = headers
.get("retry-after")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u64>().ok())
{
let reset = secs_ceil(verdict.reset_after);
if reset > current {
if let Ok(v) = reset.to_string().parse() {
headers.insert("retry-after", v);
}
}
}
}
}
#[cfg(feature = "redis")]
fn global_reject(limit: u64, window: Duration) -> Response {
let backoff = fleet_backoff(window);
let verdict = Verdict {
allowed: false,
new_tat: Duration::ZERO,
remaining: 0,
retry_after: backoff,
reset_after: backoff,
};
too_many_requests(limit, 0, &verdict)
}
struct RuleKey {
store: String,
identity: String,
}
fn rule_key(
fingerprint: &str,
key: &KeySource,
client: &str,
headers: &HeaderMap,
claims: Option<&serde_json::Value>,
) -> RuleKey {
let (tag, identity) = match key {
KeySource::Ip => ("ip", client.to_string()),
KeySource::Header(name) => match header_str(headers, name) {
Some(v) => ("hdr", v),
None => ("ip", client.to_string()),
},
KeySource::JwtClaim(claim) => match claims.and_then(|c| resolve::claim_str(c, claim)) {
Some(v) => ("jwt", v),
None => ("ip", client.to_string()),
},
};
RuleKey {
store: format!("{fingerprint}:{tag}:{}", matcher::short_hash(&identity)),
identity,
}
}
fn parse_cidr(s: &str) -> Result<ipnet::IpNet, String> {
if let Ok(net) = s.parse::<ipnet::IpNet>() {
return Ok(net);
}
if let Ok(ip) = s.parse::<std::net::IpAddr>() {
let prefix = if ip.is_ipv4() { 32 } else { 128 };
return ipnet::IpNet::new(ip, prefix)
.map_err(|e| format!("invalid trusted_proxies entry {s:?}: {e}"));
}
Err(format!("invalid trusted_proxies CIDR/IP: {s:?}"))
}
fn client_ip(
peer: Option<std::net::IpAddr>,
headers: &HeaderMap,
trusted: &[ipnet::IpNet],
) -> String {
match peer {
Some(ip) => {
if trusted.iter().any(|net| net.contains(&ip)) {
if let Some(client) = rightmost_untrusted(headers, trusted) {
return client;
}
}
ip.to_string()
}
None => "unknown".to_string(),
}
}
fn rightmost_untrusted(headers: &HeaderMap, trusted: &[ipnet::IpNet]) -> Option<String> {
if let Some(xff) = headers.get("x-forwarded-for").and_then(|v| v.to_str().ok()) {
for hop in xff.split(',').rev() {
let hop = hop.trim();
if hop.is_empty() {
continue;
}
let trusted_hop = hop
.parse::<std::net::IpAddr>()
.is_ok_and(|ip| trusted.iter().any(|net| net.contains(&ip)));
if !trusted_hop {
return Some(hop.to_string());
}
}
}
header_str(headers, "x-real-ip")
}
fn header_str(headers: &HeaderMap, name: &str) -> Option<String> {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_string)
}
fn attach_rate_headers(headers: &mut HeaderMap, limit: u64, remaining: u64, verdict: &Verdict) {
if let Ok(v) = limit.to_string().parse() {
headers.insert("ratelimit-limit", v);
}
if let Ok(v) = remaining.to_string().parse() {
headers.insert("ratelimit-remaining", v);
}
if let Ok(v) = secs_ceil(verdict.reset_after).to_string().parse() {
headers.insert("ratelimit-reset", v);
}
}
fn too_many_requests(limit: u64, remaining: u64, verdict: &Verdict) -> Response {
let mut response = (
StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({
"error": "RESOURCE_EXHAUSTED",
"message": "rate limit exceeded",
})),
)
.into_response();
let headers = response.headers_mut();
attach_rate_headers(headers, limit, remaining, verdict);
if let Ok(v) = secs_ceil(verdict.retry_after).to_string().parse() {
headers.insert("retry-after", v);
}
response
}
fn secs_ceil(d: Duration) -> u64 {
let nanos = d.as_nanos();
u64::try_from(nanos.div_ceil(1_000_000_000)).unwrap_or(u64::MAX)
}
#[cfg(test)]
mod tests;