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