use super::*;
use crate::config::{KeySourceConfig, LimitProfileConfig, RateRuleConfig, ShieldConfig};
use axum::body::Body;
use axum::http::{Request, StatusCode};
use axum::routing::get;
use axum::Router;
use std::collections::HashMap;
use tower::ServiceExt;
fn config(profiles: Vec<(&str, &str, Option<u64>)>, rules: Vec<RateRuleConfig>) -> ShieldConfig {
let profiles = profiles
.into_iter()
.map(|(name, rate, burst)| {
(
name.to_string(),
LimitProfileConfig {
rate: rate.to_string(),
burst,
},
)
})
.collect::<HashMap<_, _>>();
ShieldConfig {
enabled: true,
profiles,
rules,
default_profile: None,
jwt_limits: None,
limit_service: None,
sync: None,
trusted_proxies: Vec::new(),
}
}
fn rule(pattern: &str, key: KeySourceConfig, profile: Option<&str>) -> RateRuleConfig {
RateRuleConfig {
pattern: pattern.to_string(),
key,
profile: profile.map(str::to_string),
}
}
fn app(cfg: ShieldConfig) -> Router {
let shield = Shield::build(&cfg).unwrap().unwrap();
Router::new()
.route("/api/x", get(|| async { "ok" }))
.route("/open", get(|| async { "ok" }))
.layer(axum::middleware::from_fn_with_state(
shield,
pre_auth_middleware,
))
}
async fn get_req(app: &Router, path: &str, xff: Option<&str>) -> Response {
let mut req = Request::builder().uri(path);
if let Some(xff) = xff {
req = req.header("x-forwarded-for", xff);
}
app.clone()
.oneshot(req.body(Body::empty()).unwrap())
.await
.unwrap()
}
async fn get_keyed(app: &Router, path: &str, header: (&str, &str)) -> Response {
let req = Request::builder()
.uri(path)
.header(header.0, header.1)
.body(Body::empty())
.unwrap();
app.clone().oneshot(req).await.unwrap()
}
#[tokio::test]
async fn admits_burst_then_rejects() {
let app = app(config(
vec![("t", "60/min", Some(2))], vec![rule("/api/**", KeySourceConfig::Ip, Some("t"))],
));
assert_eq!(
get_req(&app, "/api/x", Some("9.9.9.9")).await.status(),
StatusCode::OK
);
assert_eq!(
get_req(&app, "/api/x", Some("9.9.9.9")).await.status(),
StatusCode::OK
);
let third = get_req(&app, "/api/x", Some("9.9.9.9")).await;
assert_eq!(third.status(), StatusCode::TOO_MANY_REQUESTS);
assert!(third.headers().contains_key("retry-after"));
}
#[tokio::test]
async fn emits_ratelimit_headers_when_allowed() {
let app = app(config(
vec![("t", "100/min", Some(10))],
vec![rule("/api/**", KeySourceConfig::Ip, Some("t"))],
));
let resp = get_req(&app, "/api/x", Some("1.2.3.4")).await;
assert_eq!(resp.status(), StatusCode::OK);
let h = resp.headers();
assert_eq!(h.get("ratelimit-limit").unwrap(), "100");
assert_eq!(h.get("ratelimit-remaining").unwrap(), "9");
assert!(h.contains_key("ratelimit-reset"));
}
#[tokio::test]
async fn unmatched_path_is_not_limited() {
let app = app(config(
vec![("t", "1/min", Some(1))],
vec![rule("/api/**", KeySourceConfig::Ip, Some("t"))],
));
for _ in 0..5 {
assert_eq!(
get_req(&app, "/open", Some("5.5.5.5")).await.status(),
StatusCode::OK
);
}
}
#[tokio::test]
async fn header_key_isolates_clients() {
let app = app(config(
vec![("t", "60/min", Some(1))], vec![rule(
"/api/**",
KeySourceConfig::Header {
name: "x-api-key".to_string(),
},
Some("t"),
)],
));
assert_eq!(
get_keyed(&app, "/api/x", ("x-api-key", "alice"))
.await
.status(),
StatusCode::OK
);
assert_eq!(
get_keyed(&app, "/api/x", ("x-api-key", "alice"))
.await
.status(),
StatusCode::TOO_MANY_REQUESTS
);
assert_eq!(
get_keyed(&app, "/api/x", ("x-api-key", "bob"))
.await
.status(),
StatusCode::OK
);
}
#[tokio::test]
async fn rule_without_resolvable_profile_passes() {
let app = app(config(
vec![("t", "1/min", Some(1))],
vec![rule("/api/**", KeySourceConfig::Ip, None)],
));
for _ in 0..3 {
assert_eq!(
get_req(&app, "/api/x", Some("7.7.7.7")).await.status(),
StatusCode::OK
);
}
}
#[test]
fn jwt_claim_key_uses_claim_then_falls_back_to_ip() {
let claims = serde_json::json!({ "sub": "alice", "org": { "id": "acme" } });
let key = KeySource::JwtClaim("sub".to_string());
let k = rule_key("fp", &key, "1.1.1.1", &HeaderMap::new(), Some(&claims));
assert_eq!(k.identity, "alice");
assert!(k.store.starts_with("fp:jwt:"));
assert!(
!k.store.contains("alice"),
"raw value must not appear in store key"
);
let nested = KeySource::JwtClaim("org.id".to_string());
assert_eq!(
rule_key("fp", &nested, "1.1.1.1", &HeaderMap::new(), Some(&claims)).identity,
"acme"
);
let anon = rule_key("fp", &key, "1.1.1.1", &HeaderMap::new(), None);
assert_eq!(anon.identity, "1.1.1.1");
assert!(anon.store.starts_with("fp:ip:"));
let missing = rule_key(
"fp",
&KeySource::JwtClaim("missing".to_string()),
"1.1.1.1",
&HeaderMap::new(),
Some(&claims),
);
assert_eq!(missing.identity, "1.1.1.1");
assert!(missing.store.starts_with("fp:ip:"));
}
#[test]
fn store_key_namespaced_by_fingerprint_and_value() {
let key = KeySource::Ip;
let a = rule_key("fpA", &key, "1.1.1.1", &HeaderMap::new(), None);
let b = rule_key("fpB", &key, "1.1.1.1", &HeaderMap::new(), None);
assert_ne!(a.store, b.store);
let a2 = rule_key("fpA", &key, "1.1.1.1", &HeaderMap::new(), None);
assert_eq!(a.store, a2.store);
let c = rule_key("fpA", &key, "2.2.2.2", &HeaderMap::new(), None);
assert_ne!(a.store, c.store);
}
#[test]
fn secs_ceil_rounds_sub_second_waits_up() {
use std::time::Duration;
assert_eq!(secs_ceil(Duration::from_micros(500)), 1);
assert_eq!(secs_ceil(Duration::from_millis(1)), 1);
assert_eq!(secs_ceil(Duration::from_millis(1500)), 2);
assert_eq!(secs_ceil(Duration::from_secs(3)), 3);
assert_eq!(secs_ceil(Duration::ZERO), 0);
}
#[test]
fn reconciled_headers_widen_reset_when_fleet_binds() {
use crate::shield::gcra::Verdict;
use std::time::Duration;
let verdict = Verdict {
allowed: true,
new_tat: Duration::from_millis(100),
remaining: 599,
retry_after: Duration::ZERO,
reset_after: Duration::from_millis(100),
};
let window = Duration::from_secs(60);
let (reported, hv) = reconciled_headers(verdict, Some(1), window);
assert_eq!(reported, 0, "fleet budget must bind the reported remaining");
assert!(
hv.reset_after >= Duration::from_secs(6),
"expected fleet-derived reset, got {:?}",
hv.reset_after
);
}
#[test]
fn reconciled_headers_keep_local_reset_when_local_binds() {
use crate::shield::gcra::Verdict;
use std::time::Duration;
let verdict = Verdict {
allowed: true,
new_tat: Duration::from_secs(30),
remaining: 0,
retry_after: Duration::ZERO,
reset_after: Duration::from_secs(30),
};
let (reported, hv) = reconciled_headers(verdict, Some(100), Duration::from_secs(60));
assert_eq!(reported, 0);
assert_eq!(hv.reset_after, Duration::from_secs(30));
}
#[test]
fn tighten_headers_widen_retry_after_when_outer_reset_wins() {
use crate::shield::gcra::Verdict;
use axum::http::HeaderMap;
use std::time::Duration;
let mut headers = HeaderMap::new();
headers.insert("ratelimit-remaining", "0".parse().unwrap());
headers.insert("ratelimit-reset", "60".parse().unwrap());
headers.insert("retry-after", "60".parse().unwrap());
let outer = Verdict {
allowed: false,
new_tat: Duration::ZERO,
remaining: 0,
retry_after: Duration::from_secs(3600),
reset_after: Duration::from_secs(3600),
};
maybe_tighten_rate_headers(&mut headers, 1, 0, &outer);
assert_eq!(headers.get("ratelimit-reset").unwrap(), "3600");
assert_eq!(headers.get("retry-after").unwrap(), "3600");
}
#[test]
fn build_rejects_unknown_default_profile() {
let mut cfg = config(vec![("t", "1/min", None)], Vec::new());
cfg.rules = vec![rule("/x", KeySourceConfig::Ip, Some("t"))];
cfg.default_profile = Some("missing".to_string());
assert!(Shield::build(&cfg).is_err());
}
mod two_phase {
use super::*;
use crate::auth::Auth;
use crate::config::{AuthConfig, JwtConfig};
use ed25519_dalek::{Signer, SigningKey};
fn keypair_pem() -> (SigningKey, std::path::PathBuf) {
let sk = SigningKey::from_bytes(&[9u8; 32]);
let spki_prefix: [u8; 12] = [
0x30, 0x2a, 0x30, 0x05, 0x06, 0x03, 0x2b, 0x65, 0x70, 0x03, 0x21, 0x00,
];
let mut der = spki_prefix.to_vec();
der.extend_from_slice(sk.verifying_key().as_bytes());
use base64::Engine;
let b64 = base64::engine::general_purpose::STANDARD.encode(&der);
let pem = format!("-----BEGIN PUBLIC KEY-----\n{b64}\n-----END PUBLIC KEY-----\n");
use std::sync::atomic::{AtomicU32, Ordering};
static N: AtomicU32 = AtomicU32::new(0);
let path = std::env::temp_dir().join(format!(
"sp_shield_{}_{}.pem",
std::process::id(),
N.fetch_add(1, Ordering::Relaxed)
));
std::fs::write(&path, pem).unwrap();
(sk, path)
}
fn token(sk: &SigningKey, sub: &str) -> String {
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine;
let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"EdDSA","typ":"JWT"}"#);
let claims = serde_json::json!({ "sub": sub, "exp": 9_999_999_999u64 });
let payload = URL_SAFE_NO_PAD.encode(serde_json::to_vec(&claims).unwrap());
let signing_input = format!("{header}.{payload}");
let sig = sk.sign(signing_input.as_bytes());
format!("{signing_input}.{}", URL_SAFE_NO_PAD.encode(sig.to_bytes()))
}
fn auth(pem: std::path::PathBuf) -> std::sync::Arc<Auth> {
let cfg = AuthConfig {
mode: "jwt".into(),
jwt: Some(JwtConfig {
issuer: None,
audience: None,
jwks_uri: None,
public_key_pem_file: Some(pem),
claims_headers: HashMap::new(),
roles_claim: "roles".into(),
}),
forward_auth: None,
authz: None,
};
Auth::build(&cfg).unwrap().unwrap()
}
fn stack(shield: std::sync::Arc<Shield>, auth: std::sync::Arc<Auth>) -> Router {
Router::new()
.route("/api/x", get(|| async { "ok" }))
.layer(axum::middleware::from_fn_with_state(
shield.clone(),
post_auth_middleware,
))
.layer(axum::middleware::from_fn_with_state(
auth,
crate::auth::middleware,
))
.layer(axum::middleware::from_fn_with_state(
shield,
pre_auth_middleware,
))
}
async fn get_with(app: &Router, bearer: &str) -> StatusCode {
let req = Request::builder()
.uri("/api/x")
.header("authorization", format!("Bearer {bearer}"))
.body(Body::empty())
.unwrap();
app.clone().oneshot(req).await.unwrap().status()
}
#[tokio::test]
async fn limits_per_principal_via_validated_claim() {
let (sk, pem) = keypair_pem();
let cfg = config(
vec![("t", "60/min", Some(1))], vec![rule(
"/api/**",
KeySourceConfig::JwtClaim {
claim: "sub".to_string(),
},
Some("t"),
)],
);
let shield = Shield::build(&cfg).unwrap().unwrap();
let app = stack(shield, auth(pem));
let alice = token(&sk, "alice");
let bob = token(&sk, "bob");
assert_eq!(get_with(&app, &alice).await, StatusCode::OK);
assert_eq!(get_with(&app, &alice).await, StatusCode::TOO_MANY_REQUESTS);
assert_eq!(get_with(&app, &bob).await, StatusCode::OK);
}
#[tokio::test]
async fn inner_principal_headers_survive_outer_limiter() {
let (sk, pem) = keypair_pem();
let cfg = config(
vec![("wide", "100/min", Some(100)), ("tight", "60/min", Some(1))],
vec![
rule("/api/**", KeySourceConfig::Ip, Some("wide")),
rule(
"/api/**",
KeySourceConfig::JwtClaim {
claim: "sub".to_string(),
},
Some("tight"),
),
],
);
let shield = Shield::build(&cfg).unwrap().unwrap();
let app = stack(shield, auth(pem));
let alice = token(&sk, "alice");
let send = |bearer: String| {
let app = app.clone();
async move {
let req = Request::builder()
.uri("/api/x")
.header("authorization", format!("Bearer {bearer}"))
.body(Body::empty())
.unwrap();
app.oneshot(req).await.unwrap()
}
};
assert_eq!(send(alice.clone()).await.status(), StatusCode::OK);
let rejected = send(alice).await;
assert_eq!(rejected.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(
rejected.headers().get("ratelimit-limit").unwrap(),
"60",
"the rejecting inner rule's RateLimit-Limit must not be overwritten"
);
}
#[tokio::test]
async fn allowed_response_reports_tighter_outer_budget() {
let (sk, pem) = keypair_pem();
let cfg = config(
vec![("tight", "60/min", Some(2)), ("wide", "100/min", Some(100))],
vec![
rule("/api/**", KeySourceConfig::Ip, Some("tight")),
rule(
"/api/**",
KeySourceConfig::JwtClaim {
claim: "sub".to_string(),
},
Some("wide"),
),
],
);
let shield = Shield::build(&cfg).unwrap().unwrap();
let app = stack(shield, auth(pem));
let req = Request::builder()
.uri("/api/x")
.header("authorization", format!("Bearer {}", token(&sk, "alice")))
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(resp.headers().get("ratelimit-remaining").unwrap(), "1");
}
#[tokio::test]
async fn allowed_response_reports_longer_reset_on_remaining_tie() {
let (sk, pem) = keypair_pem();
let cfg = config(
vec![
("hourly", "1/hour", Some(1)),
("minutely", "1/min", Some(1)),
],
vec![
rule("/api/**", KeySourceConfig::Ip, Some("hourly")),
rule(
"/api/**",
KeySourceConfig::JwtClaim {
claim: "sub".to_string(),
},
Some("minutely"),
),
],
);
let shield = Shield::build(&cfg).unwrap().unwrap();
let app = stack(shield, auth(pem));
let req = Request::builder()
.uri("/api/x")
.header("authorization", format!("Bearer {}", token(&sk, "alice")))
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(resp.headers().get("ratelimit-remaining").unwrap(), "0");
let reset: u64 = resp
.headers()
.get("ratelimit-reset")
.unwrap()
.to_str()
.unwrap()
.parse()
.unwrap();
assert!(reset > 120, "expected the hourly reset to win, got {reset}");
}
}