use async_trait::async_trait;
use jsonwebtoken::{decode, decode_header, jwk::JwkSet, Algorithm, DecodingKey, Validation};
use jsonwebtoken::jwk::{AlgorithmParameters, Jwk, OctetKeyParameters};
use reqwest::Client;
use std::{sync::Arc, time::Duration};
use serde::Deserialize;
use tokio::sync::RwLock;
use globset::{Glob, GlobSet, GlobSetBuilder};
use crate::{core::{ProxyError, ProxyRequest}, debug_fmt, error_fmt, security::{SecurityProvider, SecurityStage}, trace_fmt, warn_fmt};
pub const CLAIMS_ATTRIBUTE: &str = "oidc-claims";
const BEARER: &str = "bearer ";
const JWKS_REFRESH: Duration = Duration::from_secs(30 * 60);
#[derive(Debug, Clone, serde::Deserialize)]
pub struct RouteRuleConfig {
pub methods: Vec<String>,
pub path: String,
}
#[derive(Debug)]
struct RouteRule {
methods: Vec<String>,
paths: GlobSet,
}
impl RouteRule {
fn matches(&self, method: &str, path: &str) -> bool {
let method_match = self.methods.iter().any(|m| m == "*" || m == method);
let path_match = self.paths.is_match(path);
trace_fmt!("OidcProvider", "OIDC bypass rule check: method={} path={} -> method_match={} path_match={}",
method, path, method_match, path_match);
method_match && path_match
}
}
#[derive(Debug, Clone, Deserialize)]
pub struct OidcConfig {
#[serde(rename = "issuer-uri")]
pub issuer_uri: String,
pub aud: Option<String>,
#[serde(rename = "shared-secret")]
pub shared_secret: Option<String>,
#[serde(default)]
pub bypass: Vec<RouteRuleConfig>,
}
#[derive(Debug)]
pub struct OidcProvider {
issuer: String,
aud: Option<String>,
shared_secret: Option<String>,
jwks_uri: String,
jwks: Arc<RwLock<Option<JwkSet>>>,
last_refresh: Arc<RwLock<tokio::time::Instant>>,
http: Client,
rules: Vec<RouteRule>,
}
impl OidcProvider {
pub async fn discover(cfg: OidcConfig) -> Result<Self, ProxyError> {
debug_fmt!("OidcProvider", "OIDC discovery from {}", cfg.issuer_uri);
let client = Client::builder()
.user_agent("foxy/oidc")
.build()
.map_err(|e| {
let err = ProxyError::SecurityError(format!("Failed to build HTTP client: {}", e));
error_fmt!("OidcProvider", "{}", err);
err
})?;
#[derive(Deserialize)]
struct Discovery { jwks_uri: String }
let meta: Discovery = match client.get(&cfg.issuer_uri).send().await {
Ok(response) => {
match response.error_for_status() {
Ok(response) => {
match response.json().await {
Ok(meta) => meta,
Err(e) => {
let err = ProxyError::SecurityError(
format!("Failed to parse OIDC discovery response: {}", e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
},
Err(e) => {
let err = ProxyError::SecurityError(
format!("OIDC discovery endpoint returned error: {}", e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
},
Err(e) => {
let err = ProxyError::SecurityError(
format!("Failed to connect to OIDC discovery endpoint: {}", e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
};
debug_fmt!("OidcProvider", "OIDC discovery successful, JWKS URI: {}", meta.jwks_uri);
let mut rules = Vec::with_capacity(cfg.bypass.len());
for raw in cfg.bypass {
let mut builder = GlobSetBuilder::new();
match Glob::new(&raw.path) {
Ok(glob) => {
builder.add(glob);
rules.push(RouteRule {
methods: raw.methods.iter().map(|m| m.to_ascii_uppercase()).collect(),
paths: match builder.build() {
Ok(set) => set,
Err(e) => {
let err = ProxyError::SecurityError(
format!("Failed to build glob set for path {}: {}", raw.path, e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
},
});
debug_fmt!("OidcProvider", "Added OIDC bypass rule: methods={:?}, path={}", raw.methods, raw.path);
},
Err(e) => {
let err = ProxyError::SecurityError(
format!("Invalid glob pattern in bypass rule: {}", e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
}
Ok(Self {
issuer: cfg
.issuer_uri
.trim_end_matches("/.well-known/openid-configuration")
.to_owned(),
aud: cfg.aud,
shared_secret: cfg.shared_secret,
jwks_uri: meta.jwks_uri,
jwks: Arc::new(RwLock::new(None)),
last_refresh: Arc::new(RwLock::new(
tokio::time::Instant::now() - JWKS_REFRESH,
)),
http: client,
rules,
})
}
async fn refresh_jwks(&self) -> Result<(), ProxyError> {
let now = tokio::time::Instant::now();
if now.duration_since(*self.last_refresh.read().await) < JWKS_REFRESH {
trace_fmt!("OidcProvider", "JWKS cache still fresh, skipping refresh");
return Ok(());
}
debug_fmt!("OidcProvider", "Refreshing JWKS from {}", self.jwks_uri);
let jwks = match self.http.get(&self.jwks_uri).send().await {
Ok(response) => {
match response.error_for_status() {
Ok(response) => {
match response.json::<JwkSet>().await {
Ok(jwks) => jwks,
Err(e) => {
let err = ProxyError::SecurityError(
format!("Failed to parse JWKS response: {}", e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
},
Err(e) => {
let err = ProxyError::SecurityError(
format!("JWKS endpoint returned error: {}", e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
},
Err(e) => {
let err = ProxyError::SecurityError(
format!("Failed to connect to JWKS endpoint: {}", e)
);
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
};
debug_fmt!("OidcProvider", "JWKS refresh successful, found {} keys", jwks.keys.len());
{
let mut w = self.jwks.write().await;
*w = Some(jwks);
}
{
let mut w = self.last_refresh.write().await;
*w = now;
}
Ok(())
}
fn jwk_to_decoding_key(&self, jwk: &Jwk) -> Result<DecodingKey, ProxyError> {
match &jwk.algorithm {
AlgorithmParameters::RSA(params) => {
trace_fmt!("OidcProvider", "Converting RSA JWK to decoding key");
DecodingKey::from_rsa_components(¶ms.n, ¶ms.e)
.map_err(|e| {
let err = ProxyError::SecurityError(format!("Invalid RSA key: {}", e));
error_fmt!("OidcProvider", "{}", err);
err
})
}
AlgorithmParameters::EllipticCurve(params) => {
trace_fmt!("OidcProvider", "Converting EC JWK to decoding key");
DecodingKey::from_ec_components(¶ms.x, ¶ms.y)
.map_err(|e| {
let err = ProxyError::SecurityError(format!("Invalid EC key: {}", e));
error_fmt!("OidcProvider", "{}", err);
err
})
}
AlgorithmParameters::OctetKey(OctetKeyParameters { value, .. }) => {
trace_fmt!("OidcProvider", "Converting octet JWK to decoding key");
Ok(DecodingKey::from_secret(value.as_bytes()))
}
AlgorithmParameters::OctetKeyPair(params) => {
trace_fmt!("OidcProvider", "Converting OKP JWK to decoding key");
DecodingKey::from_ed_components(¶ms.x)
.map_err(|e| {
let err = ProxyError::SecurityError(format!("Invalid OKP key: {}", e));
error_fmt!("OidcProvider", "{}", err);
err
})
}
}
}
async fn validate_token(&self, token: &str) -> Result<serde_json::Value, ProxyError> {
let header = match decode_header(token) {
Ok(h) => h,
Err(e) => {
let err = ProxyError::SecurityError(format!("Invalid JWT header: {}", e));
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
};
trace_fmt!("OidcProvider", "JWT header: alg={:?}, kid={:?}", header.alg, header.kid);
let allowed_algs = [
Algorithm::RS256, Algorithm::RS384, Algorithm::RS512,
Algorithm::PS256, Algorithm::PS384, Algorithm::PS512,
Algorithm::ES256, Algorithm::ES384,
Algorithm::EdDSA,
Algorithm::HS256, Algorithm::HS384, Algorithm::HS512,
];
if !allowed_algs.contains(&header.alg) {
let err = ProxyError::SecurityError(
format!("Algorithm not allowed: {:?}", header.alg)
);
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
self.refresh_jwks().await?;
let key = match &header.kid {
Some(kid) => {
let jwks = self.jwks.read().await;
let jwks = match &*jwks {
Some(j) => j,
None => {
let err = ProxyError::SecurityError("No JWKS available".to_string());
error_fmt!("OidcProvider", "{}", err);
return Err(err);
}
};
match jwks.keys.iter().find(|k| k.common.key_id == Some(kid.clone())) {
Some(key) => {
trace_fmt!("OidcProvider", "Found key with ID {}", kid);
match self.jwk_to_decoding_key(key) {
Ok(key) => key,
Err(e) => {
error_fmt!("OidcProvider", "Failed to convert JWK to decoding key: {}", e);
return Err(e);
}
}
}
None => {
if let Some(ref secret) = self.shared_secret {
if matches!(header.alg, Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512) {
trace_fmt!("OidcProvider", "Using shared secret for HS* algorithm");
DecodingKey::from_secret(secret.as_bytes())
} else {
let err = ProxyError::SecurityError(format!("Key ID {} not found in JWKS", kid));
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
} else {
let err = ProxyError::SecurityError(format!("Key ID {} not found in JWKS", kid));
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
}
}
None => {
if let Some(ref secret) = self.shared_secret {
trace_fmt!("OidcProvider", "No key ID in token, using shared secret");
DecodingKey::from_secret(secret.as_bytes())
} else {
let err = ProxyError::SecurityError("No key ID in token and no shared secret configured".to_string());
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
};
let mut validation = Validation::new(header.alg);
validation.set_audience(&[&self.aud.clone().unwrap_or_default()]);
validation.set_issuer(&[&self.issuer]);
match decode::<serde_json::Value>(token, &key, &validation) {
Ok(token_data) => {
debug_fmt!("OidcProvider", "JWT validation successful");
Ok(token_data.claims)
}
Err(e) => {
let err = ProxyError::SecurityError(format!("JWT validation failed: {}", e));
warn_fmt!("OidcProvider", "{}", err);
Err(err)
}
}
}
fn validate_std_claims(&self, claims: &serde_json::Value) -> Result<(), ProxyError> {
if let Some(iss) = claims["iss"].as_str() {
if iss != self.issuer {
let err = ProxyError::SecurityError(
format!("Invalid issuer: expected '{}', got '{}'", self.issuer, iss)
);
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
} else {
let err = ProxyError::SecurityError("Missing issuer claim".to_string());
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
if let Some(ref expected_aud) = self.aud {
let valid_audience = match &claims["aud"] {
serde_json::Value::String(aud) => aud == expected_aud,
serde_json::Value::Array(auds) => auds.iter()
.filter_map(|a| a.as_str())
.any(|a| a == expected_aud),
_ => false,
};
if !valid_audience {
let err = ProxyError::SecurityError(
format!("Invalid audience: expected '{}'", expected_aud)
);
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
if let Some(exp) = claims["exp"].as_i64() {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64;
if exp <= now {
let err = ProxyError::SecurityError(
format!("Token expired at {}, current time is {}", exp, now)
);
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
}
debug_fmt!("OidcProvider", "Token claims validation successful");
Ok(())
}
#[inline]
fn is_bypassed(&self, method: &str, path: &str) -> bool {
let bypassed = self.rules.iter().any(|r| r.matches(method, path));
if bypassed {
debug_fmt!("OidcProvider", "OIDC bypass for {} {}", method, path);
}
bypassed
}
}
#[async_trait]
impl SecurityProvider for OidcProvider {
fn name(&self) -> &str { "OidcProvider" }
fn stage(&self) -> SecurityStage { SecurityStage::Pre }
async fn pre(&self, req: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
if self.is_bypassed(&req.method.to_string(), &req.path) {
debug_fmt!("OidcProvider", "OIDC bypass for {} {}", req.method, req.path);
return Ok(req);
}
debug_fmt!("OidcProvider", "OIDC validating request: {} {}", req.method, req.path);
let auth_header = match req.headers.get("authorization") {
Some(h) => match h.to_str() {
Ok(s) => s.to_lowercase(),
Err(e) => {
let err = ProxyError::SecurityError(
format!("Invalid authorization header: {}", e)
);
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
},
None => {
let err = ProxyError::SecurityError("Missing authorization header".to_string());
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
};
if !auth_header.starts_with(BEARER) {
let err = ProxyError::SecurityError(
format!("Invalid authorization scheme: expected 'Bearer', got '{}'",
auth_header.split_whitespace().next().unwrap_or(""))
);
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
let token = &auth_header[BEARER.len()..];
if token.is_empty() {
let err = ProxyError::SecurityError("Empty bearer token".to_string());
warn_fmt!("OidcProvider", "{}", err);
return Err(err);
}
trace_fmt!("OidcProvider", "Validating token: {}", token);
let claims = match self.validate_token(token).await {
Ok(claims) => claims,
Err(e) => {
warn_fmt!("OidcProvider", "Token validation failed: {}", e);
return Err(e);
}
};
{
let mut ctx = req.context.write().await;
ctx.attributes.insert(CLAIMS_ATTRIBUTE.to_string(), claims);
}
debug_fmt!("OidcProvider", "OIDC validation successful for {} {}", req.method, req.path);
Ok(req)
}
}