foxy/security/
oidc.rs

1// This Source Code Form is subject to the terms of the Mozilla Public
2// License, v. 2.0. If a copy of the MPL was not distributed with this
3// file, You can obtain one at https://mozilla.org/MPL/2.0/.
4
5//! OpenID-Connect bearer-token provider.
6//!
7//! Supported algs   : HS256 / 384 / 512  · RS256 / 384 / 512 · PS256 / 384 / 512
8//!                    ES256 / 384        · EdDSA (Ed25519)
9//! Bypass rules     : glob-style paths + method list, evaluated before token checks
10//! HMAC secret      : optional `shared-secret` in config (required for HS* algs)
11//! JWKS refresh     : lazy + every 30 min ± key-rotation retry
12
13use async_trait::async_trait;
14use jsonwebtoken::{decode, decode_header, jwk::JwkSet, Algorithm, DecodingKey, Validation};
15use jsonwebtoken::jwk::{AlgorithmParameters, Jwk, OctetKeyParameters};
16use reqwest::Client;
17use std::{sync::Arc, time::Duration};
18use serde::Deserialize;
19use tokio::sync::RwLock;
20use globset::{Glob, GlobSet, GlobSetBuilder};
21use crate::{core::{ProxyError, ProxyRequest}, debug_fmt, error_fmt, security::{SecurityProvider, SecurityStage}, trace_fmt, warn_fmt};
22
23pub const CLAIMS_ATTRIBUTE: &str = "oidc-claims";
24const BEARER: &str = "bearer ";
25const JWKS_REFRESH: Duration = Duration::from_secs(30 * 60);
26
27#[derive(Debug, Clone, serde::Deserialize)]
28pub struct RouteRuleConfig {
29    pub methods: Vec<String>,
30    pub path: String,
31}
32
33#[derive(Debug)]
34struct RouteRule {
35    methods: Vec<String>,
36    paths: GlobSet,
37}
38
39impl RouteRule {
40    fn matches(&self, method: &str, path: &str) -> bool {
41        let method_match = self.methods.iter().any(|m| m == "*" || m == method);
42        let path_match = self.paths.is_match(path);
43        
44        trace_fmt!("OidcProvider", "OIDC bypass rule check: method={} path={} -> method_match={} path_match={}", 
45            method, path, method_match, path_match);
46            
47        method_match && path_match
48    }
49}
50
51/// Top-level OIDC section under `"security_chain"` in config.
52#[derive(Debug, Clone, Deserialize)]
53pub struct OidcConfig {
54    #[serde(rename = "issuer-uri")]
55    pub issuer_uri: String,
56    
57    /// Expected audience claim (optional)
58    pub aud: Option<String>,
59    
60    /// Shared secret for HS* algorithms (optional)
61    #[serde(rename = "shared-secret")]
62    pub shared_secret: Option<String>,
63    
64    /// Routes to bypass authentication for
65    #[serde(default)]
66    pub bypass: Vec<RouteRuleConfig>,
67}
68
69/// OpenID Connect security provider.
70#[derive(Debug)]
71pub struct OidcProvider {
72    /// Issuer URI
73    issuer: String,
74    
75    /// Expected audience claim
76    aud: Option<String>,
77    
78    /// Shared secret for HS* algorithms
79    shared_secret: Option<String>,
80    
81    /// JWKS URI
82    jwks_uri: String,
83    
84    /// Cached JWKS
85    jwks: Arc<RwLock<Option<JwkSet>>>,
86    
87    /// Last refresh time
88    last_refresh: Arc<RwLock<tokio::time::Instant>>,
89    
90    /// HTTP client
91    http: Client,
92    
93    /// Bypass rules
94    rules: Vec<RouteRule>,
95}
96
97impl OidcProvider {
98    /// Discover OIDC configuration from the issuer URI.
99    pub async fn discover(cfg: OidcConfig) -> Result<Self, ProxyError> {
100        // --- minimal discovery ---
101        debug_fmt!("OidcProvider", "OIDC discovery from {}", cfg.issuer_uri);
102        
103        let client = Client::builder()
104            .user_agent("foxy/oidc")
105            .build()
106            .map_err(|e| {
107                let err = ProxyError::SecurityError(format!("Failed to build HTTP client: {}", e));
108                error_fmt!("OidcProvider", "{}", err);
109                err
110            })?;
111            
112        #[derive(Deserialize)]
113        struct Discovery { jwks_uri: String }
114        
115        let meta: Discovery = match client.get(&cfg.issuer_uri).send().await {
116            Ok(response) => {
117                match response.error_for_status() {
118                    Ok(response) => {
119                        match response.json().await {
120                            Ok(meta) => meta,
121                            Err(e) => {
122                                let err = ProxyError::SecurityError(
123                                    format!("Failed to parse OIDC discovery response: {}", e)
124                                );
125                                error_fmt!("OidcProvider", "{}", err);
126                                return Err(err);
127                            }
128                        }
129                    },
130                    Err(e) => {
131                        let err = ProxyError::SecurityError(
132                            format!("OIDC discovery endpoint returned error: {}", e)
133                        );
134                        error_fmt!("OidcProvider", "{}", err);
135                        return Err(err);
136                    }
137                }
138            },
139            Err(e) => {
140                let err = ProxyError::SecurityError(
141                    format!("Failed to connect to OIDC discovery endpoint: {}", e)
142                );
143                error_fmt!("OidcProvider", "{}", err);
144                return Err(err);
145            }
146        };
147
148        debug_fmt!("OidcProvider", "OIDC discovery successful, JWKS URI: {}", meta.jwks_uri);
149
150        // --- compile bypass rules ---
151        let mut rules = Vec::with_capacity(cfg.bypass.len());
152        for raw in cfg.bypass {
153            let mut builder = GlobSetBuilder::new();
154            match Glob::new(&raw.path) {
155                Ok(glob) => {
156                    builder.add(glob);
157                    rules.push(RouteRule {
158                        methods: raw.methods.iter().map(|m| m.to_ascii_uppercase()).collect(),
159                        paths: match builder.build() {
160                            Ok(set) => set,
161                            Err(e) => {
162                                let err = ProxyError::SecurityError(
163                                    format!("Failed to build glob set for path {}: {}", raw.path, e)
164                                );
165                                error_fmt!("OidcProvider", "{}", err);
166                                return Err(err);
167                            }
168                        },
169                    });
170                    debug_fmt!("OidcProvider", "Added OIDC bypass rule: methods={:?}, path={}", raw.methods, raw.path);
171                },
172                Err(e) => {
173                    let err = ProxyError::SecurityError(
174                        format!("Invalid glob pattern in bypass rule: {}", e)
175                    );
176                    error_fmt!("OidcProvider", "{}", err);
177                    return Err(err);
178                }
179            }
180        }
181
182        Ok(Self {
183            issuer: cfg
184                .issuer_uri
185                .trim_end_matches("/.well-known/openid-configuration")
186                .to_owned(),
187            aud: cfg.aud,
188            shared_secret: cfg.shared_secret,
189            jwks_uri: meta.jwks_uri,
190            jwks: Arc::new(RwLock::new(None)),
191            last_refresh: Arc::new(RwLock::new(
192                tokio::time::Instant::now() - JWKS_REFRESH,
193            )),
194            http: client,
195            rules,
196        })
197    }
198
199    /* ---------- helpers -------------------------------------------------- */
200
201    async fn refresh_jwks(&self) -> Result<(), ProxyError> {
202        let now = tokio::time::Instant::now();
203        if now.duration_since(*self.last_refresh.read().await) < JWKS_REFRESH {
204            trace_fmt!("OidcProvider", "JWKS cache still fresh, skipping refresh");
205            return Ok(());
206        }
207        
208        debug_fmt!("OidcProvider", "Refreshing JWKS from {}", self.jwks_uri);
209        
210        // Fetch the JWKS
211        let jwks = match self.http.get(&self.jwks_uri).send().await {
212            Ok(response) => {
213                match response.error_for_status() {
214                    Ok(response) => {
215                        match response.json::<JwkSet>().await {
216                            Ok(jwks) => jwks,
217                            Err(e) => {
218                                let err = ProxyError::SecurityError(
219                                    format!("Failed to parse JWKS response: {}", e)
220                                );
221                                error_fmt!("OidcProvider", "{}", err);
222                                return Err(err);
223                            }
224                        }
225                    },
226                    Err(e) => {
227                        let err = ProxyError::SecurityError(
228                            format!("JWKS endpoint returned error: {}", e)
229                        );
230                        error_fmt!("OidcProvider", "{}", err);
231                        return Err(err);
232                    }
233                }
234            },
235            Err(e) => {
236                let err = ProxyError::SecurityError(
237                    format!("Failed to connect to JWKS endpoint: {}", e)
238                );
239                error_fmt!("OidcProvider", "{}", err);
240                return Err(err);
241            }
242        };
243        
244        debug_fmt!("OidcProvider", "JWKS refresh successful, found {} keys", jwks.keys.len());
245        
246        // Update the cache
247        {
248            let mut w = self.jwks.write().await;
249            *w = Some(jwks);
250        }
251        {
252            let mut w = self.last_refresh.write().await;
253            *w = now;
254        }
255        
256        Ok(())
257    }
258
259    fn jwk_to_decoding_key(&self, jwk: &Jwk) -> Result<DecodingKey, ProxyError> {
260        match &jwk.algorithm {
261            AlgorithmParameters::RSA(params) => {
262                trace_fmt!("OidcProvider", "Converting RSA JWK to decoding key");
263                DecodingKey::from_rsa_components(&params.n, &params.e)
264                    .map_err(|e| {
265                        let err = ProxyError::SecurityError(format!("Invalid RSA key: {}", e));
266                        error_fmt!("OidcProvider", "{}", err);
267                        err
268                    })
269            }
270            AlgorithmParameters::EllipticCurve(params) => {
271                trace_fmt!("OidcProvider", "Converting EC JWK to decoding key");
272                DecodingKey::from_ec_components(&params.x, &params.y)
273                    .map_err(|e| {
274                        let err = ProxyError::SecurityError(format!("Invalid EC key: {}", e));
275                        error_fmt!("OidcProvider", "{}", err);
276                        err
277                    })
278            }
279            AlgorithmParameters::OctetKey(OctetKeyParameters { value, .. }) => {
280                trace_fmt!("OidcProvider", "Converting octet JWK to decoding key");
281                Ok(DecodingKey::from_secret(value.as_bytes()))
282            }
283            AlgorithmParameters::OctetKeyPair(params) => {
284                trace_fmt!("OidcProvider", "Converting OKP JWK to decoding key");
285                DecodingKey::from_ed_components(&params.x)
286                    .map_err(|e| {
287                        let err = ProxyError::SecurityError(format!("Invalid OKP key: {}", e));
288                        error_fmt!("OidcProvider", "{}", err);
289                        err
290                    })
291            }
292        }
293    }
294
295    async fn validate_token(&self, token: &str) -> Result<serde_json::Value, ProxyError> {
296        // Parse the header to determine the key ID and algorithm
297        let header = match decode_header(token) {
298            Ok(h) => h,
299            Err(e) => {
300                let err = ProxyError::SecurityError(format!("Invalid JWT header: {}", e));
301                warn_fmt!("OidcProvider", "{}", err);
302                return Err(err);
303            }
304        };
305        
306        trace_fmt!("OidcProvider", "JWT header: alg={:?}, kid={:?}", header.alg, header.kid);
307        
308        // Check for allowed algorithms
309        let allowed_algs = [
310            Algorithm::RS256, Algorithm::RS384, Algorithm::RS512,
311            Algorithm::PS256, Algorithm::PS384, Algorithm::PS512,
312            Algorithm::ES256, Algorithm::ES384,
313            Algorithm::EdDSA,
314            Algorithm::HS256, Algorithm::HS384, Algorithm::HS512,
315        ];
316        
317        if !allowed_algs.contains(&header.alg) {
318            let err = ProxyError::SecurityError(
319                format!("Algorithm not allowed: {:?}", header.alg)
320            );
321            warn_fmt!("OidcProvider", "{}", err);
322            return Err(err);
323        }
324
325        // Ensure we have a fresh JWKS
326        self.refresh_jwks().await?;
327
328        // Get the key
329        let key = match &header.kid {
330            Some(kid) => {
331                // Find the key in the JWKS
332                let jwks = self.jwks.read().await;
333                let jwks = match &*jwks {
334                    Some(j) => j,
335                    None => {
336                        let err = ProxyError::SecurityError("No JWKS available".to_string());
337                        error_fmt!("OidcProvider", "{}", err);
338                        return Err(err);
339                    }
340                };
341
342                // Try to find the key by ID
343                match jwks.keys.iter().find(|k| k.common.key_id == Some(kid.clone())) {
344                    Some(key) => {
345                        trace_fmt!("OidcProvider", "Found key with ID {}", kid);
346                        match self.jwk_to_decoding_key(key) {
347                            Ok(key) => key,
348                            Err(e) => {
349                                error_fmt!("OidcProvider", "Failed to convert JWK to decoding key: {}", e);
350                                return Err(e);
351                            }
352                        }
353                    }
354                    None => {
355                        // If we have a shared secret, use that for HS* algorithms
356                        if let Some(ref secret) = self.shared_secret {
357                            if matches!(header.alg, Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512) {
358                                trace_fmt!("OidcProvider", "Using shared secret for HS* algorithm");
359                                DecodingKey::from_secret(secret.as_bytes())
360                            } else {
361                                let err = ProxyError::SecurityError(format!("Key ID {} not found in JWKS", kid));
362                                warn_fmt!("OidcProvider", "{}", err);
363                                return Err(err);
364                            }
365                        } else {
366                            let err = ProxyError::SecurityError(format!("Key ID {} not found in JWKS", kid));
367                            warn_fmt!("OidcProvider", "{}", err);
368                            return Err(err);
369                        }
370                    }
371                }
372            }
373            None => {
374                // No key ID, try to use shared secret if available
375                if let Some(ref secret) = self.shared_secret {
376                    trace_fmt!("OidcProvider", "No key ID in token, using shared secret");
377                    DecodingKey::from_secret(secret.as_bytes())
378                } else {
379                    let err = ProxyError::SecurityError("No key ID in token and no shared secret configured".to_string());
380                    warn_fmt!("OidcProvider", "{}", err);
381                    return Err(err);
382                }
383            }
384        };
385
386        // Set up validation
387        let mut validation = Validation::new(header.alg);
388        validation.set_audience(&[&self.aud.clone().unwrap_or_default()]);
389        validation.set_issuer(&[&self.issuer]);
390
391        // Validate the token
392        match decode::<serde_json::Value>(token, &key, &validation) {
393            Ok(token_data) => {
394                debug_fmt!("OidcProvider", "JWT validation successful");
395                Ok(token_data.claims)
396            }
397            Err(e) => {
398                let err = ProxyError::SecurityError(format!("JWT validation failed: {}", e));
399                warn_fmt!("OidcProvider", "{}", err);
400                Err(err)
401            }
402        }
403    }
404
405    fn validate_std_claims(&self, claims: &serde_json::Value) -> Result<(), ProxyError> {
406        // Check issuer
407        if let Some(iss) = claims["iss"].as_str() {
408            if iss != self.issuer {
409                let err = ProxyError::SecurityError(
410                    format!("Invalid issuer: expected '{}', got '{}'", self.issuer, iss)
411                );
412                warn_fmt!("OidcProvider", "{}", err);
413                return Err(err);
414            }
415        } else {
416            let err = ProxyError::SecurityError("Missing issuer claim".to_string());
417            warn_fmt!("OidcProvider", "{}", err);
418            return Err(err);
419        }
420        
421        // Check audience if configured
422        if let Some(ref expected_aud) = self.aud {
423            let valid_audience = match &claims["aud"] {
424                serde_json::Value::String(aud) => aud == expected_aud,
425                serde_json::Value::Array(auds) => auds.iter()
426                    .filter_map(|a| a.as_str())
427                    .any(|a| a == expected_aud),
428                _ => false,
429            };
430            
431            if !valid_audience {
432                let err = ProxyError::SecurityError(
433                    format!("Invalid audience: expected '{}'", expected_aud)
434                );
435                warn_fmt!("OidcProvider", "{}", err);
436                return Err(err);
437            }
438        }
439        
440        // Check expiration
441        if let Some(exp) = claims["exp"].as_i64() {
442            let now = std::time::SystemTime::now()
443                .duration_since(std::time::UNIX_EPOCH)
444                .unwrap_or_default()
445                .as_secs() as i64;
446                
447            if exp <= now {
448                let err = ProxyError::SecurityError(
449                    format!("Token expired at {}, current time is {}", exp, now)
450                );
451                warn_fmt!("OidcProvider", "{}", err);
452                return Err(err);
453            }
454        }
455        
456        debug_fmt!("OidcProvider", "Token claims validation successful");
457        Ok(())
458    }
459
460    #[inline]
461    fn is_bypassed(&self, method: &str, path: &str) -> bool {
462        let bypassed = self.rules.iter().any(|r| r.matches(method, path));
463        if bypassed {
464            debug_fmt!("OidcProvider", "OIDC bypass for {} {}", method, path);
465        }
466        bypassed
467    }
468}
469
470#[async_trait]
471impl SecurityProvider for OidcProvider {
472    fn name(&self) -> &str { "OidcProvider" }
473
474    fn stage(&self) -> SecurityStage { SecurityStage::Pre }
475
476    async fn pre(&self, req: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
477        // 0) Bypass?
478        if self.is_bypassed(&req.method.to_string(), &req.path) {
479            debug_fmt!("OidcProvider", "OIDC bypass for {} {}", req.method, req.path);
480            return Ok(req);
481        }
482
483        debug_fmt!("OidcProvider", "OIDC validating request: {} {}", req.method, req.path);
484
485        // 1) Extract bearer token
486        let auth_header = match req.headers.get("authorization") {
487            Some(h) => match h.to_str() {
488                Ok(s) => s.to_lowercase(),
489                Err(e) => {
490                    let err = ProxyError::SecurityError(
491                        format!("Invalid authorization header: {}", e)
492                    );
493                    warn_fmt!("OidcProvider", "{}", err);
494                    return Err(err);
495                }
496            },
497            None => {
498                let err = ProxyError::SecurityError("Missing authorization header".to_string());
499                warn_fmt!("OidcProvider", "{}", err);
500                return Err(err);
501            }
502        };
503
504        if !auth_header.starts_with(BEARER) {
505            let err = ProxyError::SecurityError(
506                format!("Invalid authorization scheme: expected 'Bearer', got '{}'", 
507                    auth_header.split_whitespace().next().unwrap_or(""))
508            );
509            warn_fmt!("OidcProvider", "{}", err);
510            return Err(err);
511        }
512
513        let token = &auth_header[BEARER.len()..];
514        if token.is_empty() {
515            let err = ProxyError::SecurityError("Empty bearer token".to_string());
516            warn_fmt!("OidcProvider", "{}", err);
517            return Err(err);
518        }
519
520        // 2) Validate the token
521        trace_fmt!("OidcProvider", "Validating token: {}", token);
522        let claims = match self.validate_token(token).await {
523            Ok(claims) => claims,
524            Err(e) => {
525                warn_fmt!("OidcProvider", "Token validation failed: {}", e);
526                return Err(e);
527            }
528        };
529
530        // 3) Store claims in request context
531        {
532            let mut ctx = req.context.write().await;
533            ctx.attributes.insert(CLAIMS_ATTRIBUTE.to_string(), claims);
534        }
535
536        debug_fmt!("OidcProvider", "OIDC validation successful for {} {}", req.method, req.path);
537        Ok(req)
538    }
539}