#[cfg(test)]
mod security_tests {
use crate::security::oidc::{JWKS_REFRESH, RouteRule, RouteRuleConfig};
use crate::{
HttpMethod, OidcConfig, OidcProvider, ProxyError, ProxyRequest, ProxyResponse,
RequestContext, SecurityChain, SecurityProvider, SecurityStage,
};
use async_trait::async_trait;
use globset::{Glob, GlobSetBuilder};
use jsonwebtoken::jwk::JwkSet;
use reqwest::Body;
use std::sync::Arc;
use tokio::sync::RwLock;
use base64::Engine as _;
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
use reqwest::header::HeaderMap;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn create_test_request(
method: HttpMethod,
path: &str,
headers: Vec<(&'static str, &'static str)>,
) -> ProxyRequest {
let mut header_map = reqwest::header::HeaderMap::new();
for (name, value) in headers {
header_map.insert(
reqwest::header::HeaderName::from_static(name),
reqwest::header::HeaderValue::from_static(value),
);
}
ProxyRequest {
method,
path: path.to_string(),
query: None,
headers: header_map,
body: Body::from(Vec::new()),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: Some("http://test.co.za".to_string()),
}
}
#[derive(Debug)]
struct MockSecurityProvider {
bypassed: bool,
}
impl MockSecurityProvider {
fn new(bypassed: bool) -> Self {
Self { bypassed }
}
}
#[async_trait]
impl SecurityProvider for MockSecurityProvider {
fn stage(&self) -> SecurityStage {
SecurityStage::Pre
}
fn name(&self) -> &'static str {
"mock-provider"
}
async fn pre(&self, request: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
if self.bypassed {
Ok(request)
} else {
Err(ProxyError::SecurityError(
"Mock authentication failure".to_string(),
))
}
}
}
#[tokio::test]
async fn test_security_chain_with_providers() {
let mut chain = SecurityChain::new();
let failing_provider = MockSecurityProvider::new(false);
chain.add(Arc::new(failing_provider));
let request = create_test_request(HttpMethod::Get, "/api/users", vec![]);
let result = chain.apply_pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Mock authentication failure"));
} else {
panic!("Expected a SecurityError");
}
}
#[tokio::test]
async fn test_security_chain_with_oidc_bypass() {
#[derive(Debug)]
struct MockOidcProviderWithBypass;
#[async_trait]
impl SecurityProvider for MockOidcProviderWithBypass {
fn stage(&self) -> SecurityStage {
SecurityStage::Pre
}
fn name(&self) -> &'static str {
"mock-oidc-with-bypass"
}
async fn pre(&self, request: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
if request.path == "/health" {
return Ok(request); }
Err(ProxyError::SecurityError(
"OIDC validation failed".to_string(),
))
}
}
let mut chain = SecurityChain::new();
chain.add(Arc::new(MockOidcProviderWithBypass));
let request_bypassed = create_test_request(HttpMethod::Get, "/health", vec![]);
assert!(chain.apply_pre(request_bypassed).await.is_ok());
let request_blocked = create_test_request(HttpMethod::Get, "/api/data", vec![]);
assert!(chain.apply_pre(request_blocked).await.is_err());
}
#[tokio::test]
async fn test_basic_auth_provider_success() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider};
let config = BasicAuthConfig {
credentials: vec!["user1:pass1".to_string(), "user2:pass2".to_string()],
bypass: vec![],
};
let _provider = BasicAuthProvider::new(config).unwrap();
let chain = SecurityChain::from_configs(vec![crate::security::ProviderConfig {
type_: "basic".to_string(),
config: serde_json::to_value(BasicAuthConfig {
credentials: vec!["user1:pass1".to_string()],
bypass: vec![],
})
.unwrap(),
}])
.await
.unwrap();
let request = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "Basic dXNlcjE6cGFzczE=")], );
assert!(chain.apply_pre(request).await.is_ok());
}
#[tokio::test]
async fn test_basic_auth_provider_failure_invalid_credentials() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider};
let config = BasicAuthConfig {
credentials: vec!["user1:pass1".to_string()],
bypass: vec![],
};
let _provider = BasicAuthProvider::new(config).unwrap();
let chain = SecurityChain::from_configs(vec![crate::security::ProviderConfig {
type_: "basic".to_string(),
config: serde_json::to_value(BasicAuthConfig {
credentials: vec!["user1:pass1".to_string()],
bypass: vec![],
})
.unwrap(),
}])
.await
.unwrap();
let request = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "Basic dXNlcjE6d3JvbmdwYXNz")], );
assert!(chain.apply_pre(request).await.is_err());
}
#[tokio::test]
async fn test_basic_auth_provider_bypass() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider, RouteRuleConfig};
let config = BasicAuthConfig {
credentials: vec!["user1:pass1".to_string()],
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/public/*".to_string(),
}],
};
let _provider = BasicAuthProvider::new(config).unwrap();
let chain = SecurityChain::from_configs(vec![crate::security::ProviderConfig {
type_: "basic".to_string(),
config: serde_json::to_value(BasicAuthConfig {
credentials: vec!["user1:pass1".to_string()],
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/public/*".to_string(),
}],
})
.unwrap(),
}])
.await
.unwrap();
let request_bypassed = create_test_request(HttpMethod::Get, "/public/data", vec![]);
assert!(chain.apply_pre(request_bypassed).await.is_ok());
let request_blocked_no_auth = create_test_request(HttpMethod::Get, "/protected", vec![]);
assert!(chain.apply_pre(request_blocked_no_auth).await.is_err());
let request_blocked_wrong_auth = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "Basic dXNlcjE6d3JvbmdwYXNz")],
);
assert!(chain.apply_pre(request_blocked_wrong_auth).await.is_err());
}
#[test]
fn test_security_stage_is_pre() {
assert!(SecurityStage::Pre.is_pre());
assert!(SecurityStage::Both.is_pre());
assert!(!SecurityStage::Post.is_pre());
}
#[test]
fn test_security_stage_is_post() {
assert!(SecurityStage::Post.is_post());
assert!(SecurityStage::Both.is_post());
assert!(!SecurityStage::Pre.is_post());
}
#[tokio::test]
async fn test_security_chain_multiple_providers() {
let mut chain = SecurityChain::new();
let passing_provider = MockSecurityProvider::new(true);
chain.add(Arc::new(passing_provider));
let another_passing_provider = MockSecurityProvider::new(true);
chain.add(Arc::new(another_passing_provider));
let request = create_test_request(HttpMethod::Get, "/api/users", vec![]);
let result = chain.apply_pre(request).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_security_chain_mixed_providers() {
let mut chain = SecurityChain::new();
let passing_provider = MockSecurityProvider::new(true);
chain.add(Arc::new(passing_provider));
let failing_provider = MockSecurityProvider::new(false);
chain.add(Arc::new(failing_provider));
let request = create_test_request(HttpMethod::Get, "/api/users", vec![]);
let result = chain.apply_pre(request).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_security_chain_apply_post() {
#[derive(Debug)]
struct MockPostSecurityProvider {
should_fail: bool,
}
#[async_trait]
impl SecurityProvider for MockPostSecurityProvider {
fn stage(&self) -> SecurityStage {
SecurityStage::Post
}
fn name(&self) -> &'static str {
"mock-post-provider"
}
async fn post(
&self,
_request: ProxyRequest,
response: ProxyResponse,
) -> Result<ProxyResponse, ProxyError> {
if self.should_fail {
Err(ProxyError::SecurityError(
"Mock post-auth failure".to_string(),
))
} else {
Ok(response)
}
}
}
let mut chain = SecurityChain::new();
let post_provider = MockPostSecurityProvider { should_fail: false };
chain.add(Arc::new(post_provider));
let request = create_test_request(HttpMethod::Get, "/api/users", vec![]);
let response = create_test_response(200, vec![]);
let result = chain.apply_post(request, response).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_security_chain_apply_post_failure() {
#[derive(Debug)]
struct MockPostSecurityProvider {
should_fail: bool,
}
#[async_trait]
impl SecurityProvider for MockPostSecurityProvider {
fn stage(&self) -> SecurityStage {
SecurityStage::Post
}
fn name(&self) -> &'static str {
"mock-post-provider"
}
async fn post(
&self,
_request: ProxyRequest,
response: ProxyResponse,
) -> Result<ProxyResponse, ProxyError> {
if self.should_fail {
Err(ProxyError::SecurityError(
"Mock post-auth failure".to_string(),
))
} else {
Ok(response)
}
}
}
let mut chain = SecurityChain::new();
let post_provider = MockPostSecurityProvider { should_fail: true };
chain.add(Arc::new(post_provider));
let request = create_test_request(HttpMethod::Get, "/api/users", vec![]);
let response = create_test_response(200, vec![]);
let result = chain.apply_post(request, response).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Mock post-auth failure"));
} else {
panic!("Expected a SecurityError");
}
}
#[tokio::test]
async fn test_security_chain_both_stage_provider() {
#[derive(Debug)]
struct MockBothStageProvider;
#[async_trait]
impl SecurityProvider for MockBothStageProvider {
fn stage(&self) -> SecurityStage {
SecurityStage::Both
}
fn name(&self) -> &'static str {
"mock-both-provider"
}
async fn pre(&self, request: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
Ok(request)
}
async fn post(
&self,
_request: ProxyRequest,
response: ProxyResponse,
) -> Result<ProxyResponse, ProxyError> {
Ok(response)
}
}
let mut chain = SecurityChain::new();
let both_provider = MockBothStageProvider;
chain.add(Arc::new(both_provider));
let request = create_test_request(HttpMethod::Get, "/api/users", vec![]);
let result = chain.apply_pre(request.clone()).await;
assert!(result.is_ok());
let response = create_test_response(200, vec![]);
let result = chain.apply_post(request, response).await;
assert!(result.is_ok());
}
fn create_test_response(
status: u16,
headers: Vec<(&'static str, &'static str)>,
) -> ProxyResponse {
let mut header_map = reqwest::header::HeaderMap::new();
for (name, value) in headers {
header_map.insert(
reqwest::header::HeaderName::from_static(name),
reqwest::header::HeaderValue::from_static(value),
);
}
ProxyResponse {
status,
headers: header_map,
body: reqwest::Body::from(Vec::new()),
context: Arc::new(RwLock::new(crate::core::ResponseContext::default())),
}
}
#[tokio::test]
async fn test_security_chain_from_configs_empty() {
let chain = SecurityChain::from_configs(vec![]).await.unwrap();
assert_eq!(chain.providers.len(), 0);
}
#[tokio::test]
async fn test_security_chain_from_configs_unknown_provider() {
use crate::security::ProviderConfig;
let configs = vec![ProviderConfig {
type_: "unknown_provider".to_string(),
config: serde_json::json!({}),
}];
let result = SecurityChain::from_configs(configs).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Unknown security provider type"));
} else {
panic!("Expected SecurityError for unknown provider");
}
}
#[tokio::test]
async fn test_register_security_provider() {
use crate::security::{ProviderConfig, SecurityChain, register_security_provider};
#[derive(Debug)]
struct CustomSecurityProvider;
#[async_trait]
impl SecurityProvider for CustomSecurityProvider {
fn stage(&self) -> SecurityStage {
SecurityStage::Pre
}
fn name(&self) -> &'static str {
"custom"
}
async fn pre(&self, request: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
Ok(request)
}
}
register_security_provider("custom_test", |_config| {
Box::pin(
async move { Ok(Arc::new(CustomSecurityProvider) as Arc<dyn SecurityProvider>) },
)
});
let configs = vec![ProviderConfig {
type_: "custom_test".to_string(),
config: serde_json::json!({}),
}];
let chain = SecurityChain::from_configs(configs).await.unwrap();
assert_eq!(chain.providers.len(), 1);
assert_eq!(chain.providers[0].name(), "custom");
}
#[tokio::test]
async fn test_security_provider_default_pre() {
#[derive(Debug)]
struct DefaultPreProvider;
#[async_trait]
impl SecurityProvider for DefaultPreProvider {
fn stage(&self) -> SecurityStage {
SecurityStage::Pre
}
fn name(&self) -> &'static str {
"default-pre"
}
}
let provider = DefaultPreProvider;
let request = create_test_request(HttpMethod::Get, "/test", vec![]);
let result = provider.pre(request).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_security_provider_default_post() {
#[derive(Debug)]
struct DefaultPostProvider;
#[async_trait]
impl SecurityProvider for DefaultPostProvider {
fn stage(&self) -> SecurityStage {
SecurityStage::Post
}
fn name(&self) -> &'static str {
"default-post"
}
}
let provider = DefaultPostProvider;
let request = create_test_request(HttpMethod::Get, "/test", vec![]);
let response = create_test_response(200, vec![]);
let result = provider.post(request, response).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_basic_auth_provider_invalid_credential_format() {
use crate::security::basic::BasicAuthConfig;
use crate::security::basic::BasicAuthProvider;
let config = BasicAuthConfig {
credentials: vec!["invalid_format".to_string()], bypass: vec![],
};
let result = BasicAuthProvider::new(config);
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid credential format"));
} else {
panic!("Expected SecurityError for invalid credential format");
}
}
#[tokio::test]
async fn test_basic_auth_provider_invalid_glob_pattern() {
use crate::security::basic::BasicAuthProvider;
use crate::security::basic::{BasicAuthConfig, RouteRuleConfig};
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "[invalid_glob".to_string(), }],
};
let result = BasicAuthProvider::new(config);
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid glob pattern"));
} else {
panic!("Expected SecurityError for invalid glob pattern");
}
}
#[tokio::test]
async fn test_basic_auth_provider_missing_auth_header() {
use crate::security::basic::BasicAuthConfig;
use crate::security::basic::BasicAuthProvider;
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request = create_test_request(HttpMethod::Get, "/protected", vec![]);
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Missing authorization header"));
} else {
panic!("Expected SecurityError for missing auth header");
}
}
#[tokio::test]
async fn test_basic_auth_provider_invalid_auth_scheme() {
use crate::security::basic::BasicAuthConfig;
use crate::security::basic::BasicAuthProvider;
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "Bearer token123")], );
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid authorization scheme"));
} else {
panic!("Expected SecurityError for invalid auth scheme");
}
}
#[tokio::test]
async fn test_basic_auth_provider_invalid_base64() {
use crate::security::basic::BasicAuthConfig;
use crate::security::basic::BasicAuthProvider;
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "Basic invalid_base64!")], );
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Failed to base64 decode credentials"));
} else {
panic!("Expected SecurityError for invalid base64");
}
}
#[tokio::test]
async fn test_basic_auth_provider_malformed_credentials() {
use crate::security::basic::BasicAuthConfig;
use crate::security::basic::BasicAuthProvider;
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "Basic dXNlcm9ubHk=")], );
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid basic auth credential format"));
} else {
panic!("Expected SecurityError for malformed credentials");
}
}
#[tokio::test]
async fn test_route_rule_wildcard_method() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider, RouteRuleConfig};
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![RouteRuleConfig {
methods: vec!["*".to_string()], path: "/public/*".to_string(),
}],
};
let provider = BasicAuthProvider::new(config).unwrap();
let get_request = create_test_request(HttpMethod::Get, "/public/data", vec![]);
assert!(provider.pre(get_request).await.is_ok());
let post_request = create_test_request(HttpMethod::Post, "/public/data", vec![]);
assert!(provider.pre(post_request).await.is_ok());
let put_request = create_test_request(HttpMethod::Put, "/public/data", vec![]);
assert!(provider.pre(put_request).await.is_ok());
let delete_request = create_test_request(HttpMethod::Delete, "/public/data", vec![]);
assert!(provider.pre(delete_request).await.is_ok());
}
#[tokio::test]
async fn test_route_rule_specific_methods() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider, RouteRuleConfig};
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string(), "POST".to_string()],
path: "/api/*".to_string(),
}],
};
let provider = BasicAuthProvider::new(config).unwrap();
let get_request = create_test_request(HttpMethod::Get, "/api/data", vec![]);
assert!(provider.pre(get_request).await.is_ok());
let post_request = create_test_request(HttpMethod::Post, "/api/data", vec![]);
assert!(provider.pre(post_request).await.is_ok());
let put_request = create_test_request(HttpMethod::Put, "/api/data", vec![]);
assert!(provider.pre(put_request).await.is_err());
let delete_request = create_test_request(HttpMethod::Delete, "/api/data", vec![]);
assert!(provider.pre(delete_request).await.is_err());
}
#[tokio::test]
async fn test_route_rule_complex_glob_patterns() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider, RouteRuleConfig};
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/api/v*/users/*/profile".to_string(), }],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request1 = create_test_request(HttpMethod::Get, "/api/v1/users/123/profile", vec![]);
assert!(provider.pre(request1).await.is_ok());
let request2 = create_test_request(HttpMethod::Get, "/api/v2/users/456/profile", vec![]);
assert!(provider.pre(request2).await.is_ok());
let request3 = create_test_request(HttpMethod::Get, "/api/v1/users/123/settings", vec![]);
assert!(provider.pre(request3).await.is_err());
let request4 = create_test_request(HttpMethod::Get, "/api/users/123/profile", vec![]);
assert!(provider.pre(request4).await.is_err());
}
#[tokio::test]
async fn test_route_rule_exact_path_match() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider, RouteRuleConfig};
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/health".to_string(), }],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request1 = create_test_request(HttpMethod::Get, "/health", vec![]);
assert!(provider.pre(request1).await.is_ok());
let request2 = create_test_request(HttpMethod::Get, "/health/check", vec![]);
assert!(provider.pre(request2).await.is_err());
let request3 = create_test_request(HttpMethod::Get, "/healthz", vec![]);
assert!(provider.pre(request3).await.is_err());
}
#[tokio::test]
async fn test_basic_auth_provider_multiple_bypass_rules() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider, RouteRuleConfig};
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![
RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/health".to_string(),
},
RouteRuleConfig {
methods: vec!["*".to_string()],
path: "/public/*".to_string(),
},
RouteRuleConfig {
methods: vec!["POST".to_string()],
path: "/webhook".to_string(),
},
],
};
let provider = BasicAuthProvider::new(config).unwrap();
let health_request = create_test_request(HttpMethod::Get, "/health", vec![]);
assert!(provider.pre(health_request).await.is_ok());
let public_request = create_test_request(HttpMethod::Post, "/public/data", vec![]);
assert!(provider.pre(public_request).await.is_ok());
let webhook_request = create_test_request(HttpMethod::Post, "/webhook", vec![]);
assert!(provider.pre(webhook_request).await.is_ok());
let protected_request = create_test_request(HttpMethod::Get, "/protected", vec![]);
assert!(provider.pre(protected_request).await.is_err());
}
#[tokio::test]
async fn test_basic_auth_provider_empty_credentials() {
use crate::security::basic::BasicAuthConfig;
use crate::security::basic::BasicAuthProvider;
let config = BasicAuthConfig {
credentials: vec![], bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "Basic dXNlcjpwYXNz")], );
let result = provider.pre(request).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_basic_auth_provider_case_sensitive_scheme() {
use crate::security::basic::BasicAuthConfig;
use crate::security::basic::BasicAuthProvider;
let config = BasicAuthConfig {
credentials: vec!["user:pass".to_string()],
bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let request = create_test_request(
HttpMethod::Get,
"/protected",
vec![("authorization", "basic dXNlcjpwYXNz")], );
let result = provider.pre(request).await;
assert!(result.is_ok()); }
#[test]
fn test_route_rule_config_deserialization() {
let json = r#"{
"methods": ["GET", "POST"],
"path": "/api/*"
}"#;
let config: RouteRuleConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.methods, vec!["GET", "POST"]);
assert_eq!(config.path, "/api/*");
}
#[test]
fn test_route_rule_matches() {
let mut builder = GlobSetBuilder::new();
builder.add(Glob::new("/api/*").unwrap());
let paths = builder.build().unwrap();
let rule = RouteRule {
methods: vec!["GET".to_string(), "POST".to_string()],
paths,
};
assert!(rule.matches("GET", "/api/users"));
assert!(rule.matches("POST", "/api/users"));
assert!(!rule.matches("DELETE", "/api/users"));
assert!(!rule.matches("GET", "/health"));
assert!(!rule.matches("get", "/api/users"));
assert!(!rule.matches("post", "/api/users"));
}
#[test]
fn test_route_rule_wildcard_methods() {
let mut builder = GlobSetBuilder::new();
builder.add(Glob::new("/health").unwrap());
let paths = builder.build().unwrap();
let rule = RouteRule {
methods: vec!["*".to_string()],
paths,
};
assert!(rule.matches("GET", "/health"));
assert!(rule.matches("POST", "/health"));
assert!(rule.matches("DELETE", "/health"));
assert!(!rule.matches("GET", "/api"));
}
#[test]
fn test_oidc_config_deserialization() {
let json = r#"{
"issuer-uri": "https://auth.example.com",
"jwks-uri": "https://auth.example.com/.well-known/jwks.json",
"aud": "my-app",
"shared-secret": "secret123",
"bypass": [
{
"methods": ["GET"],
"path": "/health"
}
]
}"#;
let config: OidcConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.issuer_uri, "https://auth.example.com");
assert_eq!(
config.jwks_uri,
"https://auth.example.com/.well-known/jwks.json"
);
assert_eq!(config.aud, Some("my-app".to_string()));
assert_eq!(config.shared_secret, Some("secret123".to_string()));
assert_eq!(config.bypass.len(), 1);
assert_eq!(config.bypass[0].methods, vec!["GET"]);
assert_eq!(config.bypass[0].path, "/health");
}
#[test]
fn test_oidc_config_minimal() {
let json = r#"{
"issuer-uri": "https://auth.example.com",
"jwks-uri": "https://auth.example.com/.well-known/jwks.json"
}"#;
let config: OidcConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.issuer_uri, "https://auth.example.com");
assert_eq!(
config.jwks_uri,
"https://auth.example.com/.well-known/jwks.json"
);
assert_eq!(config.aud, None);
assert_eq!(config.shared_secret, None);
assert!(config.bypass.is_empty());
}
#[test]
fn test_oidc_config_empty_bypass() {
let json = r#"{
"issuer-uri": "https://auth.example.com",
"jwks-uri": "https://auth.example.com/.well-known/jwks.json",
"bypass": []
}"#;
let config: OidcConfig = serde_json::from_str(json).unwrap();
assert!(config.bypass.is_empty());
}
#[tokio::test]
async fn test_oidc_provider_discover_success() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri()),
"authorization_endpoint": format!("{}/auth", mock_server.uri()),
"token_endpoint": format!("{}/token", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: Some("test-audience".to_string()),
shared_secret: Some("test-secret".to_string()),
bypass: vec![
RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/health".to_string(),
},
RouteRuleConfig {
methods: vec!["*".to_string()],
path: "/public/*".to_string(),
},
],
};
let result = OidcProvider::discover(config.clone()).await;
assert!(result.is_ok());
let provider = result.unwrap();
assert_eq!(provider.issuer, mock_server.uri());
assert_eq!(provider.aud, Some("test-audience".to_string()));
assert_eq!(provider.shared_secret, Some("test-secret".to_string()));
assert_eq!(provider.jwks_uri, format!("{}/jwks", mock_server.uri()));
assert_eq!(provider.rules.len(), 2);
assert!(provider.is_bypassed("GET", "/health"));
assert!(provider.is_bypassed("POST", "/public/api"));
assert!(!provider.is_bypassed("POST", "/private/api"));
}
#[tokio::test]
async fn test_oidc_provider_discover_success_minimal_config() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_ok());
let provider = result.unwrap();
assert_eq!(provider.issuer, mock_server.uri());
assert_eq!(provider.aud, None);
assert_eq!(provider.shared_secret, None);
assert_eq!(provider.jwks_uri, format!("{}/jwks", mock_server.uri()));
assert!(provider.rules.is_empty());
}
#[tokio::test]
async fn test_oidc_provider_discover_success_with_well_known_suffix() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/.well-known/openid-configuration"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: format!("{}/.well-known/openid-configuration", mock_server.uri()),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_ok());
let provider = result.unwrap();
assert_eq!(
provider.issuer,
format!("{}/.well-known/openid-configuration", mock_server.uri())
);
assert_eq!(provider.jwks_uri, format!("{}/jwks", mock_server.uri()));
}
#[tokio::test]
async fn test_oidc_provider_discover_invalid_url() {
let config = OidcConfig {
issuer_uri: "invalid-url".to_string(),
jwks_uri: "invalid-jwks-url".to_string(),
aud: None,
shared_secret: None,
bypass: vec![],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_oidc_provider_discover_http_error() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(404))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_ok());
let provider = result.unwrap();
let jwks_result = provider.refresh_jwks().await;
assert!(jwks_result.is_err());
if let Err(ProxyError::SecurityError(msg)) = jwks_result {
assert!(msg.contains("JWKS endpoint returned error"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_oidc_provider_discover_invalid_json() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_string("invalid json"))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_ok());
let provider = result.unwrap();
let jwks_result = provider.refresh_jwks().await;
assert!(jwks_result.is_err());
if let Err(ProxyError::SecurityError(msg)) = jwks_result {
assert!(
msg.contains("Failed to parse JWKS response as JSON")
|| msg.contains("Failed to connect to JWKS endpoint")
|| msg.contains("JWKS endpoint returned error")
);
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_oidc_provider_discover_invalid_bypass_glob() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "[invalid-glob".to_string(), }],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid glob pattern in bypass rule"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_oidc_provider_discover_complex_bypass_rules() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![
RouteRuleConfig {
methods: vec!["get".to_string(), "post".to_string()], path: "/api/v*/health".to_string(),
},
RouteRuleConfig {
methods: vec!["*".to_string()],
path: "/static/**".to_string(),
},
],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_ok());
let provider = result.unwrap();
assert_eq!(provider.rules.len(), 2);
assert!(provider.is_bypassed("GET", "/api/v1/health"));
assert!(provider.is_bypassed("POST", "/api/v2/health"));
assert!(provider.is_bypassed("DELETE", "/static/css/style.css"));
assert!(!provider.is_bypassed("GET", "/api/v1/users"));
}
#[tokio::test]
async fn test_jwks_refresh_cache_fresh() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = Some(JwkSet { keys: vec![] }); }
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now();
}
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_jwks_refresh_success() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "test-key-1",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None; }
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 1);
assert_eq!(jwks.keys[0].common.key_id, Some("test-key-1".to_string()));
}
#[tokio::test]
async fn test_jwks_refresh_http_error() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(500))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut w = provider.last_refresh.write().await;
*w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH * 2)
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_secs(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("JWKS endpoint returned error"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_jwks_refresh_invalid_json() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_string("invalid json"))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None; }
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Failed to parse JWKS response"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_jwks_refresh_connection_error() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "http://invalid-host-12345.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(
tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH * 2)
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_secs(1))
.unwrap_or_else(tokio::time::Instant::now)
}),
)),
http: reqwest::Client::new(),
rules: vec![],
};
let result = provider.refresh_jwks().await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Failed to connect to JWKS endpoint"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_invalid_header() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let result = provider.validate_token("invalid.token.here").await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid JWT header"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_unsupported_algorithm() {
use jsonwebtoken::Header;
use serde_json::json;
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let mut _header = Header::new(jsonwebtoken::Algorithm::HS256);
_header.alg = jsonwebtoken::Algorithm::HS256;
let header_json = json!({"alg": "none", "typ": "JWT"});
let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(header_json.to_string().as_bytes());
let payload_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(json!({"sub": "test"}).to_string().as_bytes());
let token = format!("{header_b64}.{payload_b64}.signature");
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid JWT header") || msg.contains("Algorithm not allowed"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_no_jwks_available() {
use jsonwebtoken::{Algorithm, Header};
use serde_json::json;
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)), last_refresh: Arc::new(RwLock::new(
tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH * 2)
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_secs(1))
.unwrap_or_else(tokio::time::Instant::now)
}),
)),
http: reqwest::Client::new(),
rules: vec![],
};
let _header = Header {
alg: Algorithm::RS256,
kid: Some("test-key".to_string()),
..Default::default()
};
let header_json = json!({
"alg": "RS256",
"typ": "JWT",
"kid": "test-key"
});
let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(header_json.to_string().as_bytes());
let payload_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(json!({"sub": "test"}).to_string().as_bytes());
let token = format!("{header_b64}.{payload_b64}.signature");
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("No JWKS available") || msg.contains("Failed to connect"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_hmac_success() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: Some("test-audience".to_string()),
shared_secret: Some("test-secret-key".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct HmacTestClaims {
iss: String,
aud: String,
sub: String,
exp: i64,
iat: i64,
}
let claims = HmacTestClaims {
iss: "https://auth.example.com".to_string(),
aud: "test-audience".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
iat: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret-key".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_ok());
let validated_claims = result.unwrap();
assert_eq!(validated_claims["iss"], "https://auth.example.com");
assert_eq!(validated_claims["aud"], "test-audience");
assert_eq!(validated_claims["sub"], "test-user");
}
#[tokio::test]
async fn test_validate_token_hmac_wrong_secret() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("wrong-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct WrongSecretClaims {
iss: String,
sub: String,
exp: i64,
}
let claims = WrongSecretClaims {
iss: "https://auth.example.com".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("correct-secret".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("JWT validation failed: InvalidSignature"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_hmac_no_shared_secret() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None, jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct NoSecretClaims {
iss: String,
sub: String,
exp: i64,
}
let claims = NoSecretClaims {
iss: "https://auth.example.com".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(
msg.contains("No key ID in token and no shared secret configured")
|| msg.contains("HMAC algorithms require shared secret configuration")
);
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_claims_wrong_issuer() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct WrongIssuerClaims {
iss: String,
sub: String,
exp: i64,
}
let claims = WrongIssuerClaims {
iss: "https://wrong-issuer.com".to_string(), sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("InvalidIssuer"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_claims_wrong_audience() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: Some("expected-audience".to_string()),
shared_secret: Some("test-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct WrongAudClaims {
iss: String,
aud: String,
sub: String,
exp: i64,
}
let claims = WrongAudClaims {
iss: "https://auth.example.com".to_string(),
aud: "wrong-audience".to_string(), sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("InvalidAudience"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_claims_expired_token() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct ExpiredClaims {
iss: String,
sub: String,
exp: i64,
}
let claims = ExpiredClaims {
iss: "https://auth.example.com".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
- 3600) as i64, };
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("ExpiredSignature"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_claims_long_expiration() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct TestClaimsLongExp {
iss: String,
sub: String,
exp: i64,
}
let claims = TestClaimsLongExp {
iss: "https://auth.example.com".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ (100 * 365 * 24 * 3600)) as i64, };
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
if result.is_err() {
println!("Long expiration test error: {result:?}");
}
assert!(result.is_ok());
let validated_claims = result.unwrap();
assert_eq!(validated_claims["iss"], "https://auth.example.com");
assert_eq!(validated_claims["sub"], "test-user");
}
#[tokio::test]
async fn test_validate_token_missing_key_id() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "test-key-1",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
use serde_json::json;
let header_json = json!({
"alg": "RS256",
"typ": "JWT"
});
let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(header_json.to_string().as_bytes());
let payload_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(json!({"sub": "test"}).to_string().as_bytes());
let token = format!("{header_b64}.{payload_b64}.signature");
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
println!("Actual error message: {msg}");
assert!(
msg.contains("No key ID in token and no shared secret configured")
|| msg.contains("requires 'kid' (key ID) header for security")
|| msg.contains("Asymmetric algorithms require 'kid'")
);
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_key_not_found() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "different-key",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
use serde_json::json;
let header_json = json!({
"alg": "RS256",
"typ": "JWT",
"kid": "missing-key"
});
let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(header_json.to_string().as_bytes());
let payload_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(json!({"sub": "test"}).to_string().as_bytes());
let token = format!("{header_b64}.{payload_b64}.signature");
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("not found in JWKS"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_integration_full_oidc_flow_with_bypass() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: Some("test-app".to_string()),
shared_secret: Some("integration-secret".to_string()),
bypass: vec![
RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/health".to_string(),
},
RouteRuleConfig {
methods: vec!["*".to_string()],
path: "/public/*".to_string(),
},
],
};
let provider = OidcProvider::discover(config).await.unwrap();
let bypass_request = ProxyRequest {
method: HttpMethod::Get,
path: "/health".to_string(),
query: None,
headers: HeaderMap::new(),
body: reqwest::Body::from(Vec::new()),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: Some("http://test.example.com".to_string()),
};
let result = provider.pre(bypass_request).await;
assert!(result.is_ok());
let auth_request = ProxyRequest {
method: HttpMethod::Post,
path: "/api/users".to_string(),
query: None,
headers: HeaderMap::new(), body: reqwest::Body::from(Vec::new()),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: Some("http://test.example.com".to_string()),
};
let result = provider.pre(auth_request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Missing authorization header"));
} else {
panic!("Expected SecurityError for missing auth header");
}
}
#[tokio::test]
async fn test_integration_full_oidc_flow_auth_failure() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: Some("test-secret".to_string()),
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let request = ProxyRequest {
method: HttpMethod::Post,
path: "/api/users".to_string(),
query: None,
headers: HeaderMap::new(),
body: reqwest::Body::from(Vec::new()),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: Some("http://test.example.com".to_string()),
};
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Missing authorization header"));
} else {
panic!("Expected SecurityError");
}
}
#[test]
fn test_oidc_provider_is_bypassed() {
let mut builder1 = GlobSetBuilder::new();
builder1.add(Glob::new("/health").unwrap());
let paths1 = builder1.build().unwrap();
let mut builder2 = GlobSetBuilder::new();
builder2.add(Glob::new("/public/*").unwrap());
let paths2 = builder2.build().unwrap();
let rules = vec![
RouteRule {
methods: vec!["GET".to_string()],
paths: paths1,
},
RouteRule {
methods: vec!["*".to_string()],
paths: paths2,
},
];
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules,
};
assert!(provider.is_bypassed("GET", "/health"));
assert!(!provider.is_bypassed("POST", "/health"));
assert!(provider.is_bypassed("GET", "/public/api"));
assert!(provider.is_bypassed("POST", "/public/api"));
assert!(!provider.is_bypassed("GET", "/private/api"));
}
#[test]
fn test_security_provider_trait_implementation() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
assert_eq!(provider.name(), "OidcProvider");
assert_eq!(provider.stage(), SecurityStage::Pre);
}
#[tokio::test]
async fn test_oidc_provider_pre_bypass() {
let mut builder = GlobSetBuilder::new();
builder.add(Glob::new("/health").unwrap());
let paths = builder.build().unwrap();
let rules = vec![RouteRule {
methods: vec!["GET".to_string()],
paths,
}];
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules,
};
let headers = HeaderMap::new();
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/health".to_string(),
query: None,
headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let result = provider.pre(request).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_oidc_provider_pre_missing_auth_header() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let headers = HeaderMap::new();
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/api/users".to_string(),
query: None,
headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert_eq!(msg, "Missing authorization header");
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_oidc_provider_pre_invalid_auth_scheme() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let mut headers = HeaderMap::new();
headers.insert("authorization", "Basic dXNlcjpwYXNz".parse().unwrap());
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/api/users".to_string(),
query: None,
headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid authorization scheme"));
assert!(msg.contains("expected 'Bearer'"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_oidc_provider_pre_empty_bearer_token() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer ".parse().unwrap());
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/api/users".to_string(),
query: None,
headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert_eq!(msg, "Empty bearer token");
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_oidc_provider_extract_bearer_token_ensure_token_isnt_altered() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
jwks_uri: "https://auth.example.com/jwks".to_string(),
aud: None,
shared_secret: None,
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let expected_token = "abcAbcABCaBcabC";
let mut headers = HeaderMap::new();
headers.insert(
"authorization",
format!("Bearer {expected_token}").parse().unwrap(),
);
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/".to_string(),
query: None,
headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let Ok(actual_token) = provider.extract_bearer_token(&request) else {
panic!("Error obtaining token");
};
assert_eq!(actual_token, expected_token);
}
#[tokio::test]
async fn test_jwk_to_decoding_key_rsa_success() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200)
.set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "rsa-key-1",
"use": "sig",
"alg": "RS256",
"n": "0vx7agoebGcQSuuPiLJXZptN9nndrQmbXEps2aiAFbWhM78LhWx4cbbfAAtVT86zwu1RK7aPFFxuhDR1L6tSoc_BJECPebWKRXjBZCiFV4n3oknjhMstn64tZ_2W-5JsGY4Hc5n9yBXArwl93lqt7_RN5w6Cf0h4QyQ5v-65YGjQR0_FDW2QvzqY368QQMicAtaSqzs8KJZgnYb9c7d0zgdAZHzu6qMQvRL5hajrn1n91CbOpbISD08qNLyrdkt-bFTWhAI4vMQFh6WeZu0fM4lFd2NcRwr3XPksINHaQ-G_xBniIqbw0Ls1jF44-csFCur-kEgU8awapJzKnqDKgw",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
provider.refresh_jwks().await.unwrap();
let jwks = provider.jwks.read().await;
let jwks = jwks.as_ref().unwrap();
let jwk = &jwks.keys[0];
let result = provider.jwk_to_decoding_key(jwk);
assert!(result.is_ok());
}
#[tokio::test]
async fn test_jwk_to_decoding_key_ec_success() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "EC",
"kid": "ec-key-1",
"use": "sig",
"alg": "ES256",
"crv": "P-256",
"x": "f83OJ3D2xF1Bg8vub9tLe1gHMzV76e8Tus9uPHvRVEU",
"y": "x_FEzRu9m36HLN_tue659LNpXW6pCyStikYjKIWI5a0"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
provider.refresh_jwks().await.unwrap();
let jwks = provider.jwks.read().await;
let jwks = jwks.as_ref().unwrap();
let jwk = &jwks.keys[0];
let result = provider.jwk_to_decoding_key(jwk);
assert!(result.is_ok());
}
#[tokio::test]
async fn test_jwk_to_decoding_key_octet_key_success() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "oct",
"kid": "hmac-key-1",
"use": "sig",
"alg": "HS256",
"k": "GawgguFyGrWKav7AX4VKUg"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
provider.refresh_jwks().await.unwrap();
let jwks = provider.jwks.read().await;
let jwks = jwks.as_ref().unwrap();
let jwk = &jwks.keys[0];
let result = provider.jwk_to_decoding_key(jwk);
assert!(result.is_ok());
}
#[tokio::test]
async fn test_jwk_to_decoding_key_okp_success() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "OKP",
"kid": "ed25519-key-1",
"use": "sig",
"alg": "EdDSA",
"crv": "Ed25519",
"x": "11qYAYKxCrfVS_7TyWQHOg7hcvPapiMlrwIaaPcHURo"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
provider.refresh_jwks().await.unwrap();
let jwks = provider.jwks.read().await;
let jwks = jwks.as_ref().unwrap();
let jwk = &jwks.keys[0];
let result = provider.jwk_to_decoding_key(jwk);
assert!(result.is_ok());
}
#[tokio::test]
async fn test_validate_token_with_kid_fallback_to_shared_secret() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(Some(JwkSet { keys: vec![] }))), last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let mut header = Header::new(Algorithm::HS256);
header.kid = Some("missing-key".to_string());
#[derive(serde::Serialize)]
struct TestClaims {
iss: String,
sub: String,
exp: i64,
}
let claims = TestClaims {
iss: "https://auth.example.com".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(
msg.contains("not found in JWKS")
|| msg.contains("potential algorithm confusion attack")
);
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_non_hmac_algorithm_with_missing_kid() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(Some(JwkSet { keys: vec![] }))), last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use serde_json::json;
let header_json = json!({
"alg": "RS256",
"typ": "JWT",
"kid": "missing-rsa-key"
});
let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(header_json.to_string().as_bytes());
let payload_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(json!({"sub": "test"}).to_string().as_bytes());
let token = format!("{header_b64}.{payload_b64}.signature");
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("not found in JWKS"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_different_algorithms() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret-key-for-hs384".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let header = Header::new(Algorithm::HS384);
#[derive(serde::Serialize)]
struct TestClaims {
iss: String,
sub: String,
exp: i64,
}
let claims = TestClaims {
iss: "https://auth.example.com".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret-key-for-hs384".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_ok());
let validated_claims = result.unwrap();
assert_eq!(validated_claims["iss"], "https://auth.example.com");
assert_eq!(validated_claims["sub"], "test-user");
}
#[tokio::test]
async fn test_validate_token_hs512_algorithm() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret-key-for-hs512-algorithm".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let header = Header::new(Algorithm::HS512);
#[derive(serde::Serialize)]
struct TestClaims {
iss: String,
sub: String,
exp: i64,
}
let claims = TestClaims {
iss: "https://auth.example.com".to_string(),
sub: "test-user".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
};
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret-key-for-hs512-algorithm".as_ref()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_ok());
let validated_claims = result.unwrap();
assert_eq!(validated_claims["iss"], "https://auth.example.com");
assert_eq!(validated_claims["sub"], "test-user");
}
#[test]
fn test_validate_std_claims_success() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: Some("test-audience".to_string()),
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"aud": "test-audience",
"sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() + 3600)
});
let result = provider.validate_std_claims(&claims);
assert!(result.is_ok());
}
#[test]
fn test_validate_std_claims_wrong_issuer() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"iss": "https://wrong-issuer.com",
"sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() + 3600)
});
let result = provider.validate_std_claims(&claims);
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid issuer"));
assert!(msg.contains("expected 'https://auth.example.com'"));
assert!(msg.contains("got 'https://wrong-issuer.com'"));
} else {
panic!("Expected SecurityError");
}
}
#[test]
fn test_validate_std_claims_missing_issuer() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() + 3600)
});
let result = provider.validate_std_claims(&claims);
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert_eq!(msg, "Missing issuer claim");
} else {
panic!("Expected SecurityError");
}
}
#[test]
fn test_validate_std_claims_audience_array_success() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: Some("test-audience".to_string()),
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"aud": ["other-audience", "test-audience", "another-audience"],
"sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() + 3600)
});
let result = provider.validate_std_claims(&claims);
assert!(result.is_ok());
}
#[test]
fn test_validate_std_claims_audience_array_failure() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: Some("test-audience".to_string()),
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"aud": ["other-audience", "wrong-audience", "another-audience"],
"sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() + 3600)
});
let result = provider.validate_std_claims(&claims);
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid audience"));
assert!(msg.contains("expected 'test-audience'"));
} else {
panic!("Expected SecurityError");
}
}
#[test]
fn test_validate_std_claims_invalid_audience_type() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: Some("test-audience".to_string()),
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"aud": 12345, "sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() + 3600)
});
let result = provider.validate_std_claims(&claims);
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid audience"));
assert!(msg.contains("expected 'test-audience'"));
} else {
panic!("Expected SecurityError");
}
}
#[test]
fn test_validate_std_claims_expired_token() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() - 3600) });
let result = provider.validate_std_claims(&claims);
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Token expired"));
} else {
panic!("Expected SecurityError");
}
}
#[test]
fn test_validate_std_claims_no_audience_configured() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None, shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"aud": "any-audience", "sub": "test-user",
"exp": (std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs() + 3600)
});
let result = provider.validate_std_claims(&claims);
assert!(result.is_ok());
}
#[tokio::test]
async fn test_jwks_refresh_empty_keys() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": []
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 0);
}
#[tokio::test]
async fn test_jwks_refresh_cognito_format() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kid": "1234example=",
"alg": "RS256",
"kty": "RSA",
"e": "AQAB",
"n": "1234567890",
"use": "sig"
},
{
"kid": "5678example=",
"alg": "RS256",
"kty": "RSA",
"e": "AQAB",
"n": "987654321",
"use": "sig"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 2);
assert_eq!(jwks.keys[0].common.key_id, Some("1234example=".to_string()));
assert_eq!(jwks.keys[1].common.key_id, Some("5678example=".to_string()));
}
#[tokio::test]
async fn test_jwks_refresh_with_extra_fields() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "test-key-1",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB",
"x5c": ["cert1", "cert2"],
"x5t": "thumbprint",
"x5t#S256": "sha256-thumbprint",
"custom_field": "custom_value"
}
],
"cache_max_age": 3600,
"custom_metadata": {
"provider": "test"
}
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 1);
assert_eq!(jwks.keys[0].common.key_id, Some("test-key-1".to_string()));
}
#[tokio::test]
async fn test_jwks_refresh_missing_keys_field() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"metadata": "some data",
"other_field": "value"
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Failed to parse JWKS response"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_jwks_refresh_malformed_key() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "test-key-1",
"use": "sig",
"alg": "RS256",
},
{
"kty": "RSA",
"kid": "test-key-2",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
match result {
Ok(_) => {
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert!(!jwks.keys.is_empty());
}
Err(ProxyError::SecurityError(msg)) => {
assert!(msg.contains("Failed to parse JWKS response"));
}
Err(e) => panic!("Unexpected error type: {e:?}"),
}
}
#[tokio::test]
async fn test_jwks_refresh_different_content_types() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(
serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "test-key-1",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB"
}
]
})
.to_string(),
)
.insert_header("content-type", "text/plain"),
)
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 1);
assert_eq!(jwks.keys[0].common.key_id, Some("test-key-1".to_string()));
}
#[tokio::test]
async fn test_jwks_refresh_fallback_parsing() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "fallback-key-1",
"use": "sig",
"alg": "RS256",
"n": "test-modulus-fallback",
"e": "AQAB",
"x5c": ["cert1"],
"x5t": "thumbprint",
"x5t#S256": "sha256-thumbprint",
"unknown_field": "unknown_value"
}
],
"cache_max_age": 3600,
"next_update": "2024-01-01T00:00:00Z"
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 1);
assert_eq!(
jwks.keys[0].common.key_id,
Some("fallback-key-1".to_string())
);
}
#[tokio::test]
async fn test_jwks_refresh_partial_key_failure() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "valid-key-1",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB"
},
{
"kty": "RSA",
"kid": "invalid-key-1",
"use": "sig",
"alg": "RS256"
},
{
"kty": "EC",
"kid": "valid-key-2",
"use": "sig",
"alg": "ES256",
"crv": "P-256",
"x": "f83OJ3D2xF1Bg8vub9tLe1gHMzV76e8Tus9uPHvRVEU",
"y": "x_FEzRu9m36HLN_tue659LNpXW6pCyStikYjKIWI5a0"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
match result {
Ok(_) => {
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert!(!jwks.keys.is_empty());
let key_ids: Vec<_> = jwks
.keys
.iter()
.filter_map(|k| k.common.key_id.as_ref())
.collect();
assert!(
key_ids.contains(&&"valid-key-1".to_string())
|| key_ids.contains(&&"valid-key-2".to_string())
);
}
Err(ProxyError::SecurityError(msg)) => {
assert!(msg.contains("Failed to parse JWKS response"));
println!("JWKS parsing failed as expected with malformed keys: {msg}");
}
Err(e) => panic!("Unexpected error type: {e:?}"),
}
}
#[tokio::test]
async fn test_jwks_refresh_real_world_format() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"use": "sig",
"kid": "test-key-1",
"alg": "RS256",
"n": "3TZbOj2V9n8Mml9L7djP2F_qPCP4Sk7peS45-bmQHjvHRrNMFZJ_MFWe8gVpNiovr_RLDWyDWjsXwNG6Rp9ueazrGm3YqWYdMCpd9Ba3re02MDzq4glHcoGZxWQQg_qJ0b8MnG5MdI0p4VqDLhLEbJxHZz5MBgDfME07N3Zn0Lj7ytzHPpHXrhMp3zKBPWBzZShH-JG-QDLKTmODdpZaWMRG0bWo5eyfXNkp0CWTvZgxzZ5rNHHWz4Ff-6zqMSD1x8DN5x-UEcSmpWVRu1zPNMBvqPEoaJ7-xSu4BumkEWhxLkge9Z5Y2QWDKy_D5PSabJQQ3v_G4eqWa6VCT3zZw",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(
result.is_ok(),
"JWKS refresh should succeed with real-world format"
);
let jwks = provider.jwks.read().await;
assert!(jwks.is_some(), "JWKS should be cached");
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 1, "Should have exactly one key");
let key = &jwks.keys[0];
assert_eq!(
key.common.key_id,
Some("test-key-1".to_string()),
"Key ID should match"
);
assert!(key.common.key_id.is_some(), "Key should have an ID");
match &key.algorithm {
jsonwebtoken::jwk::AlgorithmParameters::RSA(rsa_params) => {
assert!(!rsa_params.n.is_empty(), "RSA modulus should not be empty");
assert!(!rsa_params.e.is_empty(), "RSA exponent should not be empty");
assert_eq!(rsa_params.e, "AQAB", "RSA exponent should be AQAB");
}
_ => panic!("Expected RSA key parameters"),
}
let decoding_key_result = provider.jwk_to_decoding_key(key);
match decoding_key_result {
Ok(_) => {
}
Err(e) => {
println!("Key conversion failed (expected with test data): {e}");
}
}
}
#[tokio::test]
async fn test_jwks_refresh_multiple_real_keys() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"use": "sig",
"kid": "rsa-key-1",
"alg": "RS256",
"n": "3TZbOj2V9n8Mml9L7djP2F_qPCP4Sk7peS45-bmQHjvHRrNMFZJ_MFWe8gVpNiovr_RLDWyDWjsXwNG6Rp9ueazrGm3YqWYdMCpd9Ba3re02MDzq4glHcoGZxWQQg_qJ0b8MnG5MdI0p4VqDLhLEbJxHZz5MBgDfME07N3Zn0Lj7ytzHPpHXrhMp3zKBPWBzZShH-JG-QDLKTmODdpZaWMRG0bWo5eyfXNkp0CWTvZgxzZ5rNHHWz4Ff-6zqMSD1x8DN5x-UEcSmpWVRu1zPNMBvqPEoaJ7-xSu4BumkEWhxLkge9Z5Y2QWDKy_D5PSabJQQ3v_G4eqWa6VCT3zZw",
"e": "AQAB"
},
{
"kty": "EC",
"use": "sig",
"kid": "ec-key-1",
"alg": "ES256",
"crv": "P-256",
"x": "f83OJ3D2xF1Bg8vub9tLe1gHMzV76e8Tus9uPHvRVEU",
"y": "x_FEzRu9m36HLN_tue659LNpXW6pCyStikYjKIWI5a0"
},
{
"kty": "RSA",
"use": "sig",
"kid": "rsa-key-2",
"alg": "RS512",
"n": "xGKzZzOjWmeZhp7wT0T-nhnpOaZrsq7qrqAxZzu6Qk2YcxjjMRnVKySNvYgFWT7JLpinabBLVRPiehFxnvaqBw",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
{
let mut jwks_w = provider.jwks.write().await;
*jwks_w = None;
}
{
let mut refresh_w = provider.last_refresh.write().await;
*refresh_w = tokio::time::Instant::now()
.checked_sub(JWKS_REFRESH + std::time::Duration::from_secs(1))
.unwrap_or_else(|| {
tokio::time::Instant::now()
.checked_sub(std::time::Duration::from_millis(1))
.unwrap_or_else(tokio::time::Instant::now)
});
}
let result = provider.refresh_jwks().await;
assert!(
result.is_ok(),
"JWKS refresh should succeed with multiple keys"
);
let jwks = provider.jwks.read().await;
assert!(jwks.is_some(), "JWKS should be cached");
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 3, "Should have exactly three keys");
let key_ids: Vec<_> = jwks
.keys
.iter()
.filter_map(|k| k.common.key_id.as_ref())
.collect();
assert!(
key_ids.contains(&&"rsa-key-1".to_string()),
"Should contain rsa-key-1"
);
assert!(
key_ids.contains(&&"ec-key-1".to_string()),
"Should contain ec-key-1"
);
assert!(
key_ids.contains(&&"rsa-key-2".to_string()),
"Should contain rsa-key-2"
);
for key in &jwks.keys {
let decoding_key_result = provider.jwk_to_decoding_key(key);
match decoding_key_result {
Ok(_) => {
println!(
"Successfully converted key {:?} to decoding key",
key.common.key_id
);
}
Err(e) => {
println!(
"Key conversion failed for {:?} (expected with test data): {e}",
key.common.key_id
);
}
}
}
}
#[tokio::test]
async fn test_oidc_provider_with_direct_jwks_uri() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/custom-jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"use": "sig",
"kid": "direct-jwks-key",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: format!("{}/issuer", mock_server.uri()),
jwks_uri: format!("{}/custom-jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
assert_eq!(
provider.jwks_uri,
format!("{}/custom-jwks", mock_server.uri())
);
let result = provider.refresh_jwks().await;
assert!(
result.is_ok(),
"JWKS refresh should succeed with direct URI"
);
let jwks = provider.jwks.read().await;
assert!(jwks.is_some());
let jwks = jwks.as_ref().unwrap();
assert_eq!(jwks.keys.len(), 1);
assert_eq!(
jwks.keys[0].common.key_id,
Some("direct-jwks-key".to_string())
);
}
#[tokio::test]
async fn test_oidc_provider_issuer_uri_no_normalization() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": []
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: format!("{}/.well-known/openid-configuration", mock_server.uri()),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
assert_eq!(
provider.issuer,
format!("{}/.well-known/openid-configuration", mock_server.uri())
);
}
#[tokio::test]
async fn test_oidc_provider_with_audience_and_secret() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": []
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: Some("test-audience".to_string()),
shared_secret: Some("test-secret".to_string()),
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
assert_eq!(provider.aud, Some("test-audience".to_string()));
assert_eq!(provider.shared_secret, Some("test-secret".to_string()));
assert_eq!(provider.jwks_uri, format!("{}/jwks", mock_server.uri()));
}
#[tokio::test]
async fn test_oidc_provider_bypass_rules_compilation() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": []
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![
RouteRuleConfig {
methods: vec!["GET".to_string(), "POST".to_string()],
path: "/health/*".to_string(),
},
RouteRuleConfig {
methods: vec!["*".to_string()],
path: "/public".to_string(),
},
],
};
let provider = OidcProvider::discover(config).await.unwrap();
assert_eq!(provider.rules.len(), 2);
assert!(provider.is_bypassed("GET", "/health/check"));
assert!(provider.is_bypassed("POST", "/health/status"));
assert!(!provider.is_bypassed("DELETE", "/health/check")); assert!(provider.is_bypassed("GET", "/public"));
assert!(provider.is_bypassed("DELETE", "/public")); assert!(!provider.is_bypassed("GET", "/private"));
}
#[tokio::test]
async fn test_oidc_provider_invalid_bypass_rule() {
let mock_server = MockServer::start().await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "[invalid-glob".to_string(), }],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid glob pattern in bypass rule"));
} else {
panic!("Expected SecurityError with invalid glob pattern");
}
}
#[tokio::test]
async fn test_oidc_provider_http_client_build_failure() {
let config = OidcConfig {
issuer_uri: "https://auth.example.com".to_string(),
jwks_uri: "https://auth.example.com/jwks".to_string(),
aud: None,
shared_secret: None,
bypass: vec![],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_oidc_provider_fallback_instant_calculation() {
let config = OidcConfig {
issuer_uri: "https://auth.example.com".to_string(),
jwks_uri: "https://auth.example.com/jwks".to_string(),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let last_refresh = *provider.last_refresh.read().await;
let now = tokio::time::Instant::now();
assert!(last_refresh <= now);
}
#[tokio::test]
async fn test_oidc_provider_discover_glob_set_build_failure() {
let mock_server = MockServer::start().await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "[invalid-glob".to_string(), }],
};
let result = OidcProvider::discover(config).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid glob pattern in bypass rule"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_jwks_fallback_parsing_empty_keys() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "unknown_type", "kid": "test-key-1"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let result = provider.refresh_jwks().await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Failed to parse JWKS response"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_jwks_fallback_parsing_success() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "test-key-1",
"use": "sig",
"alg": "RS256",
"n": "test-modulus",
"e": "AQAB",
"extra_field": "should_be_removed" }
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_clean_jwk_oct_key_type() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "oct",
"kid": "hmac-key-1",
"use": "sig",
"alg": "HS256",
"k": "dGVzdC1zZWNyZXQ" }
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_clean_jwk_okp_key_type() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "OKP",
"kid": "ed25519-key-1",
"use": "sig",
"alg": "EdDSA",
"crv": "Ed25519",
"x": "dGVzdC14LXZhbHVl" }
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let result = provider.refresh_jwks().await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_clean_jwk_unknown_key_type() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "unknown_type", "kid": "unknown-key-1",
"use": "sig",
"alg": "UNKNOWN256"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let result = provider.refresh_jwks().await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Failed to parse JWKS response"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_none_algorithm() {
let config = OidcConfig {
issuer_uri: "https://auth.example.com".to_string(),
jwks_uri: "https://auth.example.com/jwks".to_string(),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let invalid_token = "eyJhbGciOiJOT05FIiwidHlwIjoiSldUIn0.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyfQ.";
let result = provider.validate_token(invalid_token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid JWT header") || msg.contains("Algorithm not allowed"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_hmac_kid_not_found_in_jwks() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "oct",
"kid": "different-key-id", "use": "sig",
"alg": "HS256",
"k": "dGVzdC1zZWNyZXQ"
}
]
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: Some("test-secret".to_string()),
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let header = Header {
alg: Algorithm::HS256,
kid: Some("non-existent-key".to_string()),
..Default::default()
};
let claims = serde_json::json!({
"iss": mock_server.uri(),
"sub": "test-user",
"aud": "test-audience",
"exp": (chrono::Utc::now() + chrono::Duration::hours(1)).timestamp(),
"iat": chrono::Utc::now().timestamp()
});
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("test-secret".as_bytes()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("HMAC algorithm with kid") && msg.contains("not found in JWKS"));
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_jwks_response_body_read_failure() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_string("valid json"))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let result = provider.refresh_jwks().await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_validate_token_invalid_authorization_header() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: None,
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer valid-token".parse().unwrap());
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/protected".to_string(),
query: None,
headers,
body: reqwest::Body::from(Vec::new()),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let result = provider.pre(request).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_validate_token_no_shared_secret_no_kid() {
let config = OidcConfig {
issuer_uri: "https://auth.example.com".to_string(),
jwks_uri: "https://auth.example.com/jwks".to_string(),
aud: None,
shared_secret: None, bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let header = Header {
alg: Algorithm::HS256,
kid: None, ..Default::default()
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"sub": "test-user",
"exp": (chrono::Utc::now() + chrono::Duration::hours(1)).timestamp(),
"iat": chrono::Utc::now().timestamp()
});
let token = encode(
&header,
&claims,
&EncodingKey::from_secret("dummy".as_bytes()),
)
.unwrap();
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
println!("Actual error message: {msg}");
assert!(
msg.contains("No key ID in token and no shared secret configured")
|| msg.contains("HMAC algorithms require shared secret")
);
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_validate_token_asymmetric_algorithm_no_shared_secret() {
let config = OidcConfig {
issuer_uri: "https://auth.example.com".to_string(),
jwks_uri: "https://auth.example.com/jwks".to_string(),
aud: None,
shared_secret: Some("test-secret".to_string()),
bypass: vec![],
};
let provider = OidcProvider::discover(config).await.unwrap();
let header = Header {
alg: Algorithm::RS256, kid: None, ..Default::default()
};
let claims = serde_json::json!({
"iss": "https://auth.example.com",
"sub": "test-user",
"exp": (chrono::Utc::now() + chrono::Duration::hours(1)).timestamp(),
"iat": chrono::Utc::now().timestamp()
});
let token = match encode(
&header,
&claims,
&EncodingKey::from_secret("dummy".as_bytes()),
) {
Ok(token) => token,
Err(_) => {
"eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJpc3MiOiJodHRwczovL2F1dGguZXhhbXBsZS5jb20iLCJzdWIiOiJ0ZXN0LXVzZXIiLCJleHAiOjk5OTk5OTk5OTksImlhdCI6MTAwMDAwMDAwMH0.dummy-signature".to_string()
}
};
let result = provider.validate_token(&token).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
println!("Actual error message: {msg}");
assert!(
msg.contains("requires 'kid' (key ID) header")
|| msg.contains("Asymmetric algorithms require 'kid'")
|| msg.contains("Invalid JWT header")
);
} else {
panic!("Expected SecurityError");
}
}
#[tokio::test]
async fn test_multiple_bypass_rules_overlapping() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": mock_server.uri(),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let config = OidcConfig {
issuer_uri: mock_server.uri(),
jwks_uri: format!("{}/jwks", mock_server.uri()),
aud: None,
shared_secret: None,
bypass: vec![
RouteRuleConfig {
methods: vec!["GET".to_string()],
path: "/api/*".to_string(),
},
RouteRuleConfig {
methods: vec!["*".to_string()],
path: "/api/health".to_string(), },
RouteRuleConfig {
methods: vec!["POST".to_string(), "PUT".to_string()],
path: "/api/users/*".to_string(),
},
],
};
let provider = OidcProvider::discover(config).await.unwrap();
assert_eq!(provider.rules.len(), 3);
assert!(provider.is_bypassed("GET", "/api/health")); assert!(provider.is_bypassed("DELETE", "/api/health")); assert!(provider.is_bypassed("POST", "/api/users/123")); assert!(provider.is_bypassed("GET", "/api/users/123")); }
#[tokio::test]
async fn test_authorization_header_with_extra_whitespace() {
let provider = OidcProvider {
issuer: "https://auth.example.com".to_string(),
aud: None,
shared_secret: Some("test-secret".to_string()),
jwks_uri: "https://auth.example.com/jwks".to_string(),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
let mut headers = HeaderMap::new();
headers.insert("authorization", " Bearer token123 ".parse().unwrap());
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/api/users".to_string(),
query: None,
headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let result = provider.pre(request).await;
assert!(result.is_err());
if let Err(ProxyError::SecurityError(msg)) = result {
assert!(msg.contains("Invalid authorization scheme"));
} else {
panic!("Expected SecurityError");
}
}
#[test]
fn test_route_rule_matches_edge_cases() {
let mut builder = GlobSetBuilder::new();
builder.add(Glob::new("/**").unwrap()); let paths = builder.build().unwrap();
let rule = RouteRule {
methods: vec!["GET".to_string(), "POST".to_string()],
paths,
};
assert!(rule.matches("GET", "/"));
assert!(rule.matches("POST", "/api"));
assert!(rule.matches("GET", "/api/v1/users/123"));
assert!(rule.matches("POST", "/very/deep/nested/path/structure"));
assert!(!rule.matches("DELETE", "/api")); assert!(!rule.matches("PUT", "/")); }
#[test]
fn test_route_rule_matches_empty_methods() {
let mut builder = GlobSetBuilder::new();
builder.add(Glob::new("/health").unwrap());
let paths = builder.build().unwrap();
let rule = RouteRule {
methods: vec![], paths,
};
assert!(!rule.matches("GET", "/health"));
assert!(!rule.matches("POST", "/health"));
assert!(!rule.matches("*", "/health"));
}
#[tokio::test]
async fn test_jwt_algorithm_confusion_attack() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/.well-known/openid_configuration"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": format!("{}", mock_server.uri()),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
let rsa_public_key_n = "0vx7agoebGcQSuuPiLJXZptN9nndrQmbXEps2aiAFbWhM78LhWx4cbbfAAtVT86zwu1RK7aPFFxuhDR1L6tSoc_BJECPebWKRXjBZCiFV4n3oknjhMstn64tZ_2W-5JsGY4Hc5n9yBXArwl93lqt7_RN5w6Cf0h4QyQ5v-65YGjQR0_FDW2QvzqY368QQMicAtaSqzs8KJZgnYb9c7d0zgdAZHzu6qMQvRL5hajrn1n91CbOpbISD08qNLyrdkt-bFTWhAI4vMQFh6WeZu0fM4lFd2NcRwr3XPksINHaQ-G_xBniIqbw0Ls1jF44-csFCur-kEgU8awapJzKnqDKgw";
let rsa_public_key_e = "AQAB";
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "rsa-key-2022",
"use": "sig",
"alg": "RS256",
"n": rsa_public_key_n,
"e": rsa_public_key_e
}
]
})))
.mount(&mock_server)
.await;
let provider = OidcProvider {
issuer: mock_server.uri().to_string(),
aud: Some("vulnerable-api".to_string()),
shared_secret: None, jwks_uri: format!("{}/jwks", mock_server.uri()),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use base64::Engine as _;
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let rsa_modulus_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(rsa_public_key_n)
.expect("Failed to decode RSA modulus");
let mut malicious_header = Header::new(Algorithm::HS256);
malicious_header.kid = Some("rsa-key-2022".to_string());
#[derive(serde::Serialize)]
struct MaliciousClaims {
iss: String,
aud: String,
sub: String,
exp: i64,
iat: i64,
role: String, }
let malicious_claims = MaliciousClaims {
iss: mock_server.uri().to_string(),
aud: "vulnerable-api".to_string(),
sub: "attacker".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
iat: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64,
role: "admin".to_string(), };
let malicious_token = encode(
&malicious_header,
&malicious_claims,
&EncodingKey::from_secret(&rsa_modulus_bytes),
)
.expect("Failed to create malicious token");
println!("Created malicious JWT token with algorithm confusion attack");
println!("Token header algorithm: HS256 (but using RSA key as HMAC secret)");
println!("Token kid: rsa-key-2022 (points to RSA key in JWKS)");
let result = provider.validate_token(&malicious_token).await;
match result {
Ok(_) => {
panic!(
"CRITICAL SECURITY VULNERABILITY: Algorithm confusion attack succeeded! \
The OIDC provider accepted a malicious JWT token created using algorithm confusion. \
This allows complete authentication bypass and privilege escalation."
);
}
Err(e) => {
println!("✓ Algorithm confusion attack properly rejected: {e}");
let error_msg = e.to_string();
assert!(
error_msg.contains("Algorithm not allowed")
|| error_msg.contains("Invalid")
|| error_msg.contains("validation failed")
|| error_msg.contains("Key ID")
|| error_msg.contains("algorithm"),
"Token should be rejected due to algorithm/key validation, got: {error_msg}"
);
}
}
}
#[tokio::test]
async fn test_jwt_algorithm_confusion_with_shared_secret_fallback() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/.well-known/openid_configuration"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": format!("{}", mock_server.uri()),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "legitimate-key",
"use": "sig",
"alg": "RS256",
"n": "different-modulus-value",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let weak_shared_secret = "publicly-known-secret";
let provider = OidcProvider {
issuer: mock_server.uri().to_string(),
aud: Some("vulnerable-api".to_string()),
shared_secret: Some(weak_shared_secret.to_string()),
jwks_uri: format!("{}/jwks", mock_server.uri()),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
let malicious_header = Header::new(Algorithm::HS256);
#[derive(serde::Serialize)]
struct AttackClaims {
iss: String,
aud: String,
sub: String,
exp: i64,
iat: i64,
admin: bool,
}
let attack_claims = AttackClaims {
iss: mock_server.uri().to_string(),
aud: "vulnerable-api".to_string(),
sub: "attacker".to_string(),
exp: (std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
+ 3600) as i64,
iat: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64,
admin: true, };
let attack_token = encode(
&malicious_header,
&attack_claims,
&EncodingKey::from_secret(weak_shared_secret.as_ref()),
)
.expect("Failed to create attack token");
println!("Created attack token using known shared secret");
println!("Token algorithm: HS256 (no kid specified)");
let result = provider.validate_token(&attack_token).await;
match result {
Ok(claims) => {
println!("⚠️ WARNING: Attack token was accepted!");
println!("Validated claims: {claims:?}");
assert_eq!(claims["sub"], "attacker");
assert_eq!(claims["admin"], true);
println!("🚨 SECURITY ISSUE: Shared secret allows authentication bypass");
println!(" Recommendation: Use only asymmetric algorithms (RS256, ES256)");
println!(" Recommendation: Disable HMAC algorithms in production");
}
Err(e) => {
println!("✓ Attack token properly rejected: {e}");
let error_msg = e.to_string();
assert!(
error_msg.contains("Algorithm not allowed")
|| error_msg.contains("validation failed")
|| error_msg.contains("Invalid")
|| error_msg.contains("expired")
|| error_msg.contains("issuer")
|| error_msg.contains("kid"),
"Token should be rejected for security reasons, got: {error_msg}"
);
}
}
}
#[tokio::test]
async fn test_algorithm_downgrade_attack_prevention() {
let mock_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/.well-known/openid_configuration"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"issuer": format!("{}", mock_server.uri()),
"jwks_uri": format!("{}/jwks", mock_server.uri())
})))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"keys": [
{
"kty": "RSA",
"kid": "strong-key",
"use": "sig",
"alg": "RS256", "n": "strong-key-modulus",
"e": "AQAB"
}
]
})))
.mount(&mock_server)
.await;
let provider = OidcProvider {
issuer: mock_server.uri().to_string(),
aud: Some("secure-api".to_string()),
shared_secret: None, jwks_uri: format!("{}/jwks", mock_server.uri()),
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(tokio::time::Instant::now())),
http: reqwest::Client::new(),
rules: vec![],
};
use serde_json::json;
let downgrade_header = json!({
"alg": "none", "typ": "JWT",
"kid": "strong-key"
});
let payload = json!({
"iss": format!("{}", mock_server.uri()),
"aud": "secure-api",
"sub": "attacker",
"exp": 9999999999i64
});
use base64::Engine as _;
let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(downgrade_header.to_string().as_bytes());
let payload_b64 =
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(payload.to_string().as_bytes());
let unsigned_token = format!("{header_b64}.{payload_b64}.");
println!("Testing algorithm downgrade to 'none'");
let result = provider.validate_token(&unsigned_token).await;
match result {
Ok(_) => {
panic!(
"CRITICAL VULNERABILITY: 'none' algorithm was accepted! \
This allows complete authentication bypass."
);
}
Err(e) => {
println!("✓ Algorithm downgrade attack properly rejected: {e}");
let error_msg = e.to_string();
assert!(
error_msg.contains("Algorithm not allowed")
|| error_msg.contains("Invalid")
|| error_msg.contains("validation failed"),
"Should reject 'none' algorithm, got: {error_msg}"
);
}
}
}
#[cfg(feature = "vault-config")]
#[tokio::test]
async fn test_vault_path_traversal_attack() {
use crate::config::{ConfigProviderExt, VaultConfigProvider};
use std::fs;
use tempfile::tempdir;
#[derive(Debug)]
struct MockConfigProvider {
values: std::collections::HashMap<String, serde_json::Value>,
}
impl MockConfigProvider {
fn new() -> Self {
Self {
values: std::collections::HashMap::new(),
}
}
}
impl crate::config::ConfigProvider for MockConfigProvider {
fn get_raw(
&self,
key: &str,
) -> Result<Option<serde_json::Value>, crate::config::ConfigError> {
Ok(self.values.get(key).cloned())
}
fn has(&self, key: &str) -> bool {
self.values.contains_key(key)
}
fn provider_name(&self) -> &str {
"mock"
}
}
let dir = tempdir().unwrap();
let vault_dir = dir.path().join("vault");
fs::create_dir_all(&vault_dir).unwrap();
let sensitive_file = dir.path().join("sensitive.txt");
fs::write(&sensitive_file, "SENSITIVE_DATA").unwrap();
let mut mock_provider = MockConfigProvider::new();
mock_provider.values.insert(
"server.secret".to_string(),
serde_json::json!("${secret.../../sensitive}"),
);
let vault_provider = VaultConfigProvider::wrap(mock_provider, vault_dir.to_str().unwrap());
let result = vault_provider.get::<String>("server.secret");
assert!(result.is_err());
let error = result.unwrap_err();
assert!(error.to_string().contains("invalid secret name"));
println!("✓ Path traversal attack properly rejected: {}", error);
}
#[cfg(feature = "vault-config")]
#[tokio::test]
async fn test_vault_symlink_following_attack() {
use crate::config::{ConfigProviderExt, VaultConfigProvider};
use std::fs;
use tempfile::tempdir;
#[derive(Debug)]
struct MockConfigProvider {
values: std::collections::HashMap<String, serde_json::Value>,
}
impl MockConfigProvider {
fn new() -> Self {
Self {
values: std::collections::HashMap::new(),
}
}
}
impl crate::config::ConfigProvider for MockConfigProvider {
fn get_raw(
&self,
key: &str,
) -> Result<Option<serde_json::Value>, crate::config::ConfigError> {
Ok(self.values.get(key).cloned())
}
fn has(&self, key: &str) -> bool {
self.values.contains_key(key)
}
fn provider_name(&self) -> &str {
"mock"
}
}
let dir = tempdir().unwrap();
let vault_dir = dir.path().join("vault");
fs::create_dir_all(&vault_dir).unwrap();
let sensitive_file = dir.path().join("passwd");
fs::write(&sensitive_file, "root:x:0:0:root:/root:/bin/bash").unwrap();
#[cfg(unix)]
{
use std::os::unix::fs::symlink;
let symlink_path = vault_dir.join("malicious_link");
let _ = symlink(&sensitive_file, &symlink_path);
let mut mock_provider = MockConfigProvider::new();
mock_provider.values.insert(
"server.secret".to_string(),
serde_json::json!("${secret.malicious_link}"),
);
let vault_provider =
VaultConfigProvider::wrap(mock_provider, vault_dir.to_str().unwrap());
let result = vault_provider.get::<String>("server.secret");
match result {
Ok(Some(content)) => {
if content.contains("root:x:0:0") {
println!(
"WARNING: Symlink following allowed access to sensitive file: {}",
content
);
println!("This could be a security vulnerability if not intended");
}
}
Ok(None) => {
println!("✓ Secret not found (expected behavior)");
}
Err(e) => {
println!("✓ Symlink access properly rejected: {}", e);
}
}
}
#[cfg(windows)]
{
let mut mock_provider = MockConfigProvider::new();
mock_provider.values.insert(
"server.secret".to_string(),
serde_json::json!("${secret.C:\\Windows\\System32\\drivers\\etc\\hosts}"),
);
let vault_provider =
VaultConfigProvider::wrap(mock_provider, vault_dir.to_str().unwrap());
let result = vault_provider.get::<String>("server.secret");
assert!(result.is_err());
println!("✓ Absolute path properly rejected on Windows");
}
}
#[tokio::test]
async fn test_basic_auth_timing_attack_mitigation() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider};
use crate::{HttpMethod, ProxyRequest, RequestContext};
use base64::Engine as _;
use reqwest::header::HeaderMap;
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::RwLock;
let config = BasicAuthConfig {
credentials: vec![
"validuser1:validpass1".to_string(),
"validuser2:validpass2".to_string(),
"validuser3:validpass3".to_string(),
],
bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let start = Instant::now();
let result1 = provider.validate_credentials_constant_time("validuser1", "wrongpass");
let time1 = start.elapsed();
let start = Instant::now();
let result2 = provider.validate_credentials_constant_time("invaliduser", "wrongpass");
let time2 = start.elapsed();
let start = Instant::now();
let result3 = provider.validate_credentials_constant_time("validuser1", "validpass1");
let time3 = start.elapsed();
assert!(!result1); assert!(!result2); assert!(result3);
let max_diff = std::cmp::max(
time1.as_nanos().abs_diff(time2.as_nanos()),
time2.as_nanos().abs_diff(time3.as_nanos()),
);
if max_diff > 1_000_000 {
println!("WARNING: Timing difference detected: {max_diff} ns");
println!("Valid user/wrong pass: {time1:?}");
println!("Invalid user/wrong pass: {time2:?}");
println!("Valid user/valid pass: {time3:?}");
} else {
println!("✓ Constant-time comparison working - max difference: {max_diff} ns");
}
let mut headers = HeaderMap::new();
let auth_value = base64::engine::general_purpose::STANDARD.encode("validuser1:validpass1");
headers.insert(
"authorization",
format!("Basic {auth_value}").parse().unwrap(),
);
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/protected".to_string(),
query: None,
headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let result = provider.pre(request).await;
assert!(result.is_ok());
println!("✓ Basic authentication timing attack mitigation verified");
}
#[tokio::test]
async fn test_basic_auth_timing_attack_protection() {
use crate::security::basic::{BasicAuthConfig, BasicAuthProvider};
use crate::{HttpMethod, ProxyRequest, RequestContext};
use base64::Engine as _;
use reqwest::header::HeaderMap;
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::RwLock;
let config = BasicAuthConfig {
credentials: vec!["validuser:validpass".to_string()],
bypass: vec![],
};
let provider = BasicAuthProvider::new(config).unwrap();
let valid_user_creds =
base64::engine::general_purpose::STANDARD.encode("validuser:wrongpass");
let mut headers1 = HeaderMap::new();
headers1.insert(
"authorization",
format!("Basic {valid_user_creds}").parse().unwrap(),
);
let request1 = ProxyRequest {
method: HttpMethod::Get,
path: "/api/test".to_string(),
query: None,
headers: headers1,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let invalid_user_creds =
base64::engine::general_purpose::STANDARD.encode("invaliduser:wrongpass");
let mut headers2 = HeaderMap::new();
headers2.insert(
"authorization",
format!("Basic {invalid_user_creds}").parse().unwrap(),
);
let request2 = ProxyRequest {
method: HttpMethod::Get,
path: "/api/test".to_string(),
query: None,
headers: headers2,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let mut valid_user_times = Vec::new();
let mut invalid_user_times = Vec::new();
for _ in 0..10 {
let start = Instant::now();
let _ = provider.pre(request1.clone()).await;
valid_user_times.push(start.elapsed());
let start = Instant::now();
let _ = provider.pre(request2.clone()).await;
invalid_user_times.push(start.elapsed());
}
let avg_valid =
valid_user_times.iter().sum::<std::time::Duration>() / valid_user_times.len() as u32;
let avg_invalid = invalid_user_times.iter().sum::<std::time::Duration>()
/ invalid_user_times.len() as u32;
let time_diff = avg_valid.abs_diff(avg_invalid);
if time_diff.as_millis() > 1 {
println!("WARNING: Potential timing attack vulnerability detected");
println!("Average time for valid username: {avg_valid:?}");
println!("Average time for invalid username: {avg_invalid:?}");
println!("Time difference: {time_diff:?}");
}
assert!(provider.pre(request1).await.is_err());
assert!(provider.pre(request2).await.is_err());
}
#[tokio::test]
async fn test_request_smuggling_header_injection_mitigation() {
use crate::server::validate_headers;
use reqwest::header::{HeaderMap, HeaderValue};
let mut headers = HeaderMap::new();
headers.insert("content-length", "100".parse().unwrap());
headers.insert("transfer-encoding", "chunked".parse().unwrap());
let result = validate_headers(&mut headers);
assert!(result.is_ok());
assert!(!headers.contains_key("content-length"));
assert!(headers.contains_key("transfer-encoding"));
let mut headers = HeaderMap::new();
let malicious_value = "value\r\nInjected-Header: malicious";
match HeaderValue::from_str(malicious_value) {
Ok(value) => {
headers.insert("test-header", value);
let result = validate_headers(&mut headers);
assert!(result.is_err());
if let Err(e) = result {
assert!(e.to_string().contains("CRLF injection"));
}
}
Err(_) => {
println!("✓ HeaderValue::from_str properly rejects CRLF injection");
}
}
let mut headers = HeaderMap::new();
headers.append("host", "example.com".parse().unwrap());
headers.append("host", "attacker.com".parse().unwrap());
let result = validate_headers(&mut headers);
assert!(result.is_err());
if let Err(e) = result {
assert!(e.to_string().contains("Multiple Host headers"));
}
let mut headers = HeaderMap::new();
headers.insert("content-length", "100,200".parse().unwrap());
let result = validate_headers(&mut headers);
assert!(result.is_err());
if let Err(e) = result {
assert!(e.to_string().contains("Multiple Content-Length values"));
}
let mut headers = HeaderMap::new();
headers.insert("transfer-encoding", "chunked, gzip".parse().unwrap());
let result = validate_headers(&mut headers);
assert!(result.is_ok());
assert_eq!(headers.get("transfer-encoding").unwrap(), "chunked");
println!("✓ All request smuggling attack vectors properly mitigated");
}
#[tokio::test]
async fn test_request_smuggling_header_injection() {
use crate::{HttpMethod, ProxyRequest, RequestContext};
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use std::str::FromStr;
use std::sync::Arc;
use tokio::sync::RwLock;
let mut headers = HeaderMap::new();
headers.insert("content-length", "44".parse().unwrap());
headers.insert("transfer-encoding", "chunked".parse().unwrap());
headers.insert("host", "target.com".parse().unwrap());
let request = ProxyRequest {
method: HttpMethod::Post,
path: "/api/endpoint".to_string(),
query: None,
headers: headers.clone(),
body: reqwest::Body::from(
"0\r\n\r\nGET /admin/secret HTTP/1.1\r\nHost: target.com\r\n\r\n",
),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let has_content_length = request.headers.contains_key("content-length");
let has_transfer_encoding = request.headers.contains_key("transfer-encoding");
if has_content_length && has_transfer_encoding {
println!("WARNING: Request contains both Content-Length and Transfer-Encoding headers");
println!("This could enable HTTP request smuggling attacks");
}
let _malicious_headers = HeaderMap::new();
let malicious_value =
"legitimate-value\r\nX-Injected-Header: malicious\r\nX-Another-Header: attack";
match HeaderValue::from_str(malicious_value) {
Ok(_) => {
panic!("SECURITY VULNERABILITY: CRLF injection in header value was accepted");
}
Err(_) => {
println!("✓ CRLF injection in header value properly rejected");
}
}
let malicious_header_name = "X-Test\r\nX-Injected: malicious";
match HeaderName::from_str(malicious_header_name) {
Ok(_) => {
panic!("SECURITY VULNERABILITY: CRLF injection in header name was accepted");
}
Err(_) => {
println!("✓ CRLF injection in header name properly rejected");
}
}
let mut override_headers = HeaderMap::new();
override_headers.insert("x-http-method-override", "DELETE".parse().unwrap());
override_headers.insert("x-forwarded-host", "attacker.com".parse().unwrap());
let override_request = ProxyRequest {
method: HttpMethod::Post,
path: "/readonly-endpoint".to_string(),
query: None,
headers: override_headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
if override_request
.headers
.contains_key("x-http-method-override")
{
println!("WARNING: Request contains X-HTTP-Method-Override header");
println!("This could bypass method-based security controls");
}
if override_request.headers.contains_key("x-forwarded-host") {
println!("WARNING: Request contains X-Forwarded-Host header");
println!("This could enable host header injection attacks");
}
}
#[tokio::test]
async fn test_router_input_validation_mitigation() {
use crate::router::Predicate;
use crate::router::QueryPredicateConfig;
use crate::router::predicates::QueryPredicate;
use crate::{HttpMethod, ProxyRequest, RequestContext};
use reqwest::header::HeaderMap;
use std::sync::Arc;
use tokio::sync::RwLock;
let config = QueryPredicateConfig {
params: vec![("param".to_string(), "safe".to_string())]
.into_iter()
.collect(),
exact_match: false,
};
let predicate = QueryPredicate::new(config);
let crlf_query = "param=value%0d%0aInjected-Header:%20malicious";
let crlf_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(crlf_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let _result = predicate.matches(&crlf_request).await;
println!("✓ CRLF injection handled safely");
let xss_query = "param=<script>alert('xss')</script>";
let xss_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(xss_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let _result = predicate.matches(&xss_request).await;
println!("✓ XSS patterns detected and logged");
let traversal_query = "param=../../../etc/passwd";
let traversal_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(traversal_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let _result = predicate.matches(&traversal_request).await;
println!("✓ Path traversal patterns detected and logged");
let sql_query = "param='; DROP TABLE users; --";
let sql_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(sql_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let _result = predicate.matches(&sql_request).await;
println!("✓ SQL injection patterns detected and logged");
let cmd_query = "param=test; rm -rf /";
let cmd_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(cmd_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let _result = predicate.matches(&cmd_request).await;
println!("✓ Command injection patterns detected and logged");
let null_query = "param=test\0malicious";
let null_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(null_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let _result = predicate.matches(&null_request).await;
println!("✓ Null byte injection detected and sanitized");
let long_query = format!("param={}", "A".repeat(10000));
let long_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(long_query),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let _result = predicate.matches(&long_request).await;
println!("✓ Query length limits enforced");
println!("✓ All input validation mitigations verified");
}
#[tokio::test]
async fn test_router_input_validation_attacks() {
use crate::QueryPredicate;
use crate::router::Predicate;
use crate::router::QueryPredicateConfig;
use crate::{HttpMethod, ProxyRequest, RequestContext};
use reqwest::header::HeaderMap;
use std::sync::Arc;
use tokio::sync::RwLock;
let config = QueryPredicateConfig {
params: vec![("param".to_string(), "value".to_string())]
.into_iter()
.collect(),
exact_match: false,
};
let predicate = QueryPredicate::new(config);
let malicious_query = "param=value%0d%0aInjected-Header:%20malicious";
let request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(malicious_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let matches = predicate.matches(&request).await;
if matches {
println!("Query predicate matched request with potentially malicious query parameters");
if let Some(query) = &request.query {
if query.contains("\r\n") || query.contains("%0d%0a") {
println!("WARNING: Query contains CRLF sequences that could enable injection");
}
}
}
let sql_injection_query = "param='; DROP TABLE users; --";
let sql_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(sql_injection_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let sql_matches = predicate.matches(&sql_request).await;
if sql_matches {
if let Some(query) = &sql_request.query {
if query.contains("DROP") || query.contains("--") || query.contains("'") {
println!("WARNING: Query contains potential SQL injection patterns: {query}");
}
}
}
let xss_query = "param=<script>alert('xss')</script>";
let xss_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(xss_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let xss_matches = predicate.matches(&xss_request).await;
if xss_matches {
if let Some(query) = &xss_request.query {
if query.contains("<script>") || query.contains("javascript:") {
println!("WARNING: Query contains potential XSS patterns: {query}");
}
}
}
let traversal_query = "param=../../../etc/passwd";
let traversal_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api".to_string(),
query: Some(traversal_query.to_string()),
headers: HeaderMap::new(),
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let traversal_matches = predicate.matches(&traversal_request).await;
if traversal_matches {
if let Some(query) = &traversal_request.query {
if query.contains("../") || query.contains("..\\") {
println!("WARNING: Query contains potential path traversal patterns: {query}");
}
}
}
}
#[tokio::test]
async fn test_information_disclosure_in_errors() {
use crate::core::ProxyError;
let timeout_error = ProxyError::Timeout(std::time::Duration::from_secs(30));
let error_msg = timeout_error.to_string();
let sensitive_patterns = [
"/etc/passwd",
"/home/",
"C:\\Users\\",
"127.0.0.1",
"localhost",
"password",
"secret",
"key",
"token",
"internal",
"debug",
"stack trace",
];
for pattern in &sensitive_patterns {
if error_msg.to_lowercase().contains(&pattern.to_lowercase()) {
println!("WARNING: Error message may contain sensitive information: {pattern}");
println!("Error message: {error_msg}");
}
}
let routing_error =
ProxyError::RoutingError("No route found for /internal/admin/config".to_string());
let routing_msg = routing_error.to_string();
if routing_msg.contains("/internal/") || routing_msg.contains("/admin/") {
println!("WARNING: Routing error may reveal internal application structure");
println!("Error message: {routing_msg}");
}
let security_error = ProxyError::SecurityError("JWT validation failed: invalid signature from issuer https://internal.auth.company.com".to_string());
let security_msg = security_error.to_string();
if security_msg.contains("internal.") || security_msg.contains(".company.com") {
println!("WARNING: Security error may reveal internal infrastructure details");
println!("Error message: {security_msg}");
}
let config_error = ProxyError::ConfigError(
"Failed to load config from /opt/foxy/config/production.json".to_string(),
);
let config_msg = config_error.to_string();
if config_msg.contains("/opt/") || config_msg.contains("production") {
println!("WARNING: Configuration error may reveal deployment details");
println!("Error message: {config_msg}");
}
assert!(!error_msg.is_empty());
assert!(!routing_msg.is_empty());
assert!(!security_msg.is_empty());
assert!(!config_msg.is_empty());
}
#[tokio::test]
async fn test_missing_security_headers() {
use crate::core::{ProxyResponse, ResponseContext};
use reqwest::header::HeaderMap;
use std::sync::Arc;
use tokio::sync::RwLock;
let mut headers = HeaderMap::new();
headers.insert("content-type", "text/html".parse().unwrap());
headers.insert("content-length", "100".parse().unwrap());
let response = ProxyResponse {
status: 200,
headers: headers.clone(),
body: reqwest::Body::from("<!DOCTYPE html><html><body>Test</body></html>"),
context: Arc::new(RwLock::new(ResponseContext::default())),
};
let required_security_headers = [
("x-frame-options", "Security header to prevent clickjacking"),
(
"x-content-type-options",
"Security header to prevent MIME sniffing",
),
("x-xss-protection", "Security header for XSS protection"),
(
"strict-transport-security",
"Security header for HTTPS enforcement",
),
(
"content-security-policy",
"Security header to prevent XSS and injection",
),
(
"referrer-policy",
"Security header to control referrer information",
),
(
"permissions-policy",
"Security header to control browser features",
),
];
let mut missing_headers = Vec::new();
for (header_name, description) in &required_security_headers {
if !response.headers.contains_key(*header_name) {
missing_headers.push((*header_name, *description));
}
}
if !missing_headers.is_empty() {
println!("WARNING: Response is missing security headers:");
for (header, desc) in &missing_headers {
println!(" - {header}: {desc}");
}
}
if let Some(frame_options) = response.headers.get("x-frame-options") {
let value = frame_options.to_str().unwrap_or("");
if !["DENY", "SAMEORIGIN"].contains(&value) {
println!("WARNING: X-Frame-Options header has weak value: {value}");
}
}
if let Some(content_type_options) = response.headers.get("x-content-type-options") {
let value = content_type_options.to_str().unwrap_or("");
if value != "nosniff" {
println!("WARNING: X-Content-Type-Options should be 'nosniff', got: {value}");
}
}
if let Some(server_header) = response.headers.get("server") {
let value = server_header.to_str().unwrap_or("");
if value.contains("version") || value.contains("/") {
println!("WARNING: Server header may reveal version information: {value}");
}
}
assert_eq!(response.status, 200);
assert!(response.headers.contains_key("content-type"));
}
#[tokio::test]
async fn test_proxy_safeguard_bypass_techniques() {
use crate::{HttpMethod, ProxyRequest, RequestContext};
use reqwest::header::HeaderMap;
use std::sync::Arc;
use tokio::sync::RwLock;
let mut host_injection_headers = HeaderMap::new();
host_injection_headers.insert("host", "internal.service.local".parse().unwrap());
host_injection_headers.insert("x-forwarded-host", "attacker.com".parse().unwrap());
let host_injection_request = ProxyRequest {
method: HttpMethod::Get,
path: "/admin".to_string(),
query: None,
headers: host_injection_headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
if let Some(host) = host_injection_request.headers.get("host") {
if let Some(forwarded_host) = host_injection_request.headers.get("x-forwarded-host") {
println!("WARNING: Request contains both Host and X-Forwarded-Host headers");
println!("Host: {host:?}, X-Forwarded-Host: {forwarded_host:?}");
println!("This could enable host header injection attacks");
}
}
let mut method_override_headers = HeaderMap::new();
method_override_headers.insert("x-http-method-override", "DELETE".parse().unwrap());
let method_override_request = ProxyRequest {
method: HttpMethod::Post,
path: "/readonly-endpoint".to_string(),
query: None,
headers: method_override_headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
if method_override_request
.headers
.contains_key("x-http-method-override")
{
println!("WARNING: Request uses X-HTTP-Method-Override header");
println!("Original method: {:?}", method_override_request.method);
if let Some(override_method) = method_override_request
.headers
.get("x-http-method-override")
{
println!("Override method: {override_method:?}");
println!("This could bypass method-based access controls");
}
}
let mut protocol_headers = HeaderMap::new();
protocol_headers.insert("upgrade", "h2c".parse().unwrap());
protocol_headers.insert("connection", "Upgrade".parse().unwrap());
let protocol_request = ProxyRequest {
method: HttpMethod::Get,
path: "/secure-endpoint".to_string(),
query: None,
headers: protocol_headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
if let Some(upgrade) = protocol_request.headers.get("upgrade") {
if let Some(connection) = protocol_request.headers.get("connection") {
println!("WARNING: Request attempts protocol upgrade/downgrade");
println!("Upgrade: {upgrade:?}, Connection: {connection:?}");
println!("This could bypass protocol-based security controls");
}
}
let mut duplicate_headers = HeaderMap::new();
duplicate_headers.insert("authorization", "Bearer token1".parse().unwrap());
duplicate_headers.append("authorization", "Bearer token2".parse().unwrap());
let duplicate_request = ProxyRequest {
method: HttpMethod::Get,
path: "/api/secure".to_string(),
query: None,
headers: duplicate_headers,
body: reqwest::Body::from(""),
context: Arc::new(RwLock::new(RequestContext::default())),
custom_target: None,
};
let auth_headers: Vec<_> = duplicate_request
.headers
.get_all("authorization")
.iter()
.collect();
if auth_headers.len() > 1 {
println!("WARNING: Request contains multiple Authorization headers");
for (i, header) in auth_headers.iter().enumerate() {
println!(" Authorization[{i}]: {header:?}");
}
println!("This could cause inconsistent authentication behavior");
}
assert!(!host_injection_request.path.is_empty());
assert!(!method_override_request.path.is_empty());
assert!(!protocol_request.path.is_empty());
assert!(!duplicate_request.path.is_empty());
}
#[tokio::test]
async fn test_configuration_security_issues() {
use std::fs;
use tempfile::tempdir;
let dir = tempdir().unwrap();
let config_file = dir.path().join("config.json");
let insecure_config = r#"{
"database": {
"password": "hardcoded_password_123",
"connection_string": "postgresql://user:secret@localhost:5432/db"
},
"api": {
"key": "sk_live_abcd1234567890",
"secret": "very_secret_key"
},
"jwt": {
"signing_key": "super_secret_jwt_key_that_should_not_be_here"
}
}"#;
fs::write(&config_file, insecure_config).unwrap();
let config_content = fs::read_to_string(&config_file).unwrap();
let secret_patterns = [
("password", "Hardcoded password detected"),
("secret", "Hardcoded secret detected"),
("key", "Hardcoded key detected"),
("token", "Hardcoded token detected"),
("sk_live_", "Live API key detected"),
("sk_test_", "Test API key detected"),
];
let mut found_secrets = Vec::new();
for (pattern, description) in &secret_patterns {
if config_content
.to_lowercase()
.contains(&pattern.to_lowercase())
{
found_secrets.push((*pattern, *description));
}
}
if !found_secrets.is_empty() {
println!("WARNING: Configuration file contains potential secrets:");
for (pattern, desc) in &found_secrets {
println!(" - {pattern}: {desc}");
}
println!("Secrets should be externalized using environment variables or vault systems");
}
let metadata = fs::metadata(&config_file).unwrap();
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let permissions = metadata.permissions();
let mode = permissions.mode();
if mode & 0o004 != 0 {
println!("WARNING: Configuration file is world-readable");
println!("File permissions: {mode:o}");
println!("Consider restricting permissions to owner only (600)");
}
if mode & 0o040 != 0 {
println!("WARNING: Configuration file is group-readable");
println!("Consider restricting permissions to owner only (600)");
}
}
let config_size = metadata.len();
if config_size > 1024 * 1024 {
println!("WARNING: Configuration file is unusually large: {config_size} bytes");
println!("Large config files may indicate embedded secrets or excessive complexity");
}
let env_vars_to_check = [
"DATABASE_PASSWORD",
"API_KEY",
"JWT_SECRET",
"PRIVATE_KEY",
"ACCESS_TOKEN",
];
for env_var in &env_vars_to_check {
if let Ok(value) = std::env::var(env_var) {
if !value.is_empty() {
println!("INFO: Environment variable {env_var} is set");
if value.len() < 10 {
println!(
"WARNING: {} appears to have a short value ({})",
env_var,
value.len()
);
println!("Short secrets may be vulnerable to brute force attacks");
}
}
}
}
assert!(config_file.exists());
assert!(!config_content.is_empty());
}
#[tokio::test]
async fn test_dependency_security_issues() {
let vulnerable_patterns = [
("jsonwebtoken", "9", "Check for latest security patches"),
("reqwest", "0.11", "Verify TLS configuration"),
(
"hyper",
"0.14",
"Ensure latest version for HTTP/2 security fixes",
),
("tokio", "1.0", "Check for async security issues"),
("serde", "1.0", "Verify deserialization security"),
];
println!("Dependency Security Analysis:");
for (crate_name, version_pattern, recommendation) in &vulnerable_patterns {
println!(" - {crate_name}: {version_pattern} - {recommendation}");
}
let audit_findings = [
(
"RUSTSEC-2021-0124",
"jsonwebtoken",
"Algorithm confusion vulnerability",
),
(
"RUSTSEC-2022-0013",
"regex",
"ReDoS vulnerability in regex parsing",
),
("RUSTSEC-2020-0071", "time", "Segfault in time crate"),
];
let mut security_advisories = Vec::new();
for (advisory_id, crate_name, description) in &audit_findings {
security_advisories.push((*advisory_id, *crate_name, *description));
}
if !security_advisories.is_empty() {
println!("WARNING: Potential security advisories found:");
for (id, crate_name, desc) in &security_advisories {
println!(" - {id}: {crate_name} - {desc}");
}
println!("Run 'cargo audit' to check for actual vulnerabilities");
}
let internal_crates = ["foxy-internal", "company-auth", "internal-utils"];
println!("Dependency Confusion Risk Assessment:");
for crate_name in &internal_crates {
println!(
" - Check if '{crate_name}' exists on crates.io to prevent dependency confusion"
);
}
let restricted_licenses = ["GPL-3.0", "AGPL-3.0", "SSPL-1.0", "Commons Clause"];
println!("License Compliance Check:");
for license in &restricted_licenses {
println!(" - Ensure no dependencies use restricted license: {license}");
}
let supply_chain_checks = [
"Verify all dependencies are from trusted sources",
"Check for typosquatting in dependency names",
"Ensure dependency signatures are verified",
"Monitor for suspicious dependency updates",
"Use dependency pinning in production",
];
println!("Supply Chain Security Checklist:");
for check in &supply_chain_checks {
println!(" - {check}");
}
}
}