1use 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#[derive(Debug, Clone, Deserialize)]
64pub struct OidcConfig {
65 #[serde(rename = "issuer-uri")]
66 pub issuer_uri: String,
67
68 #[serde(rename = "jwks-uri")]
70 pub jwks_uri: String,
71
72 pub aud: Option<String>,
74
75 #[serde(rename = "shared-secret")]
77 pub shared_secret: Option<String>,
78
79 #[serde(default)]
81 pub bypass: Vec<RouteRuleConfig>,
82}
83
84#[derive(Debug)]
86pub struct OidcProvider {
87 pub(crate) issuer: String,
89
90 pub(crate) aud: Option<String>,
92
93 pub(crate) shared_secret: Option<String>,
95
96 pub(crate) jwks_uri: String,
98
99 pub(crate) jwks: Arc<RwLock<Option<JwkSet>>>,
101
102 pub(crate) last_refresh: Arc<RwLock<tokio::time::Instant>>,
104
105 pub(crate) http: Client,
107
108 pub(crate) rules: Vec<RouteRule>,
110}
111
112impl OidcProvider {
113 pub async fn discover(cfg: OidcConfig) -> Result<Self, ProxyError> {
115 let client = Client::builder()
116 .user_agent("foxy/oidc")
117 .build()
118 .map_err(|e| {
119 let err = ProxyError::SecurityError(format!("Failed to build HTTP client: {e}"));
120 error_fmt!("OidcProvider", "{}", err);
121 err
122 })?;
123
124 let jwks_uri = cfg.jwks_uri.clone();
126 debug_fmt!("OidcProvider", "Using JWKS URI: {}", jwks_uri);
127
128 let mut rules = Vec::with_capacity(cfg.bypass.len());
130 for raw in cfg.bypass {
131 let mut builder = GlobSetBuilder::new();
132 match Glob::new(&raw.path) {
133 Ok(glob) => {
134 builder.add(glob);
135 rules.push(RouteRule {
136 methods: raw.methods.iter().map(|m| m.to_ascii_uppercase()).collect(),
137 paths: match builder.build() {
138 Ok(set) => set,
139 Err(e) => {
140 let err = ProxyError::SecurityError(format!(
141 "Failed to build glob set for path {}: {}",
142 raw.path, e
143 ));
144 error_fmt!("OidcProvider", "{}", err);
145 return Err(err);
146 }
147 },
148 });
149 debug_fmt!(
150 "OidcProvider",
151 "Added OIDC bypass rule: methods={:?}, path={}",
152 raw.methods,
153 raw.path
154 );
155 }
156 Err(e) => {
157 let err = ProxyError::SecurityError(format!(
158 "Invalid glob pattern in bypass rule: {e}"
159 ));
160 error_fmt!("OidcProvider", "{}", err);
161 return Err(err);
162 }
163 }
164 }
165
166 Ok(Self {
167 issuer: cfg.issuer_uri,
168 aud: cfg.aud,
169 shared_secret: cfg.shared_secret,
170 jwks_uri,
171 jwks: Arc::new(RwLock::new(None)),
172 last_refresh: Arc::new(RwLock::new(
173 tokio::time::Instant::now()
174 .checked_sub(JWKS_REFRESH * 2)
175 .unwrap_or_else(|| {
176 tokio::time::Instant::now()
178 .checked_sub(std::time::Duration::from_secs(1))
179 .unwrap_or_else(tokio::time::Instant::now)
180 }),
181 )),
182 http: client,
183 rules,
184 })
185 }
186
187 fn parse_jwks_fallback(json_value: &serde_json::Value) -> Result<JwkSet, String> {
191 let keys_array = json_value
196 .get("keys")
197 .and_then(|v| v.as_array())
198 .ok_or("Missing or invalid 'keys' field")?;
199
200 let mut cleaned_keys = Vec::new();
201
202 for key_value in keys_array {
203 if let Some(cleaned_key) = Self::clean_jwk_for_parsing(key_value) {
204 cleaned_keys.push(cleaned_key);
205 }
206 }
207
208 if cleaned_keys.is_empty() {
209 return Err("No valid keys found in JWKS after cleaning".to_string());
210 }
211
212 let cleaned_jwks = serde_json::json!({
213 "keys": cleaned_keys
214 });
215
216 serde_json::from_value::<JwkSet>(cleaned_jwks)
218 .map_err(|e| format!("Failed to parse cleaned JWKS: {e}"))
219 }
220
221 fn clean_jwk_for_parsing(key_value: &serde_json::Value) -> Option<serde_json::Value> {
223 let obj = key_value.as_object()?;
224
225 let kty = obj.get("kty")?.as_str()?;
227
228 let mut cleaned = serde_json::Map::new();
229
230 for field in &["kty", "kid", "use", "alg"] {
232 if let Some(value) = obj.get(*field) {
233 cleaned.insert(field.to_string(), value.clone());
234 }
235 }
236
237 match kty {
239 "RSA" => {
240 for field in &["n", "e"] {
241 if let Some(value) = obj.get(*field) {
242 cleaned.insert(field.to_string(), value.clone());
243 }
244 }
245 }
246 "EC" => {
247 for field in &["x", "y", "crv"] {
248 if let Some(value) = obj.get(*field) {
249 cleaned.insert(field.to_string(), value.clone());
250 }
251 }
252 }
253 "oct" => {
254 if let Some(value) = obj.get("k") {
255 cleaned.insert("k".to_string(), value.clone());
256 }
257 }
258 "OKP" => {
259 for field in &["x", "crv"] {
260 if let Some(value) = obj.get(*field) {
261 cleaned.insert(field.to_string(), value.clone());
262 }
263 }
264 }
265 _ => {
266 warn_fmt!("OidcProvider", "Unknown key type in JWKS: {}", kty);
268 return None;
269 }
270 }
271
272 Some(serde_json::Value::Object(cleaned))
273 }
274
275 pub(crate) async fn refresh_jwks(&self) -> Result<(), ProxyError> {
276 let now = tokio::time::Instant::now();
277
278 let should_refresh = {
280 let jwks_guard = self.jwks.read().await;
281 let cache_empty = jwks_guard.is_none();
282 let cache_expired = now.duration_since(*self.last_refresh.read().await) >= JWKS_REFRESH;
283 cache_empty || cache_expired
284 };
285
286 if !should_refresh {
287 trace_fmt!("OidcProvider", "JWKS cache still fresh, skipping refresh");
288 return Ok(());
289 }
290
291 debug_fmt!("OidcProvider", "Refreshing JWKS from {}", self.jwks_uri);
292
293 let jwks = match self.http.get(&self.jwks_uri).send().await {
295 Ok(response) => match response.error_for_status() {
296 Ok(response) => {
297 let response_text = match response.text().await {
299 Ok(text) => text,
300 Err(e) => {
301 let err = ProxyError::SecurityError(format!(
302 "Failed to read JWKS response body: {e}"
303 ));
304 error_fmt!("OidcProvider", "{}", err);
305 return Err(err);
306 }
307 };
308
309 debug_fmt!("OidcProvider", "JWKS response body: {}", response_text);
310
311 let json_value: serde_json::Value = match serde_json::from_str(&response_text) {
313 Ok(value) => value,
314 Err(e) => {
315 let err = ProxyError::SecurityError(format!(
316 "Failed to parse JWKS response as JSON: {e}. Response body: {}",
317 response_text.chars().take(500).collect::<String>()
318 ));
319 error_fmt!("OidcProvider", "{}", err);
320 return Err(err);
321 }
322 };
323
324 match serde_json::from_value::<JwkSet>(json_value.clone()) {
326 Ok(jwks) => jwks,
327 Err(e) => {
328 debug_fmt!(
329 "OidcProvider",
330 "Standard JwkSet parsing failed: {}, trying fallback",
331 e
332 );
333
334 match Self::parse_jwks_fallback(&json_value) {
336 Ok(jwks) => {
337 debug_fmt!("OidcProvider", "Fallback JWKS parsing successful");
338 jwks
339 }
340 Err(fallback_err) => {
341 let err = ProxyError::SecurityError(format!(
342 "Failed to parse JWKS response. Standard error: {e}. Fallback error: {fallback_err}. JSON structure: {}",
343 serde_json::to_string_pretty(&json_value)
344 .unwrap_or_else(|_| "invalid".to_string())
345 ));
346 error_fmt!("OidcProvider", "{}", err);
347 return Err(err);
348 }
349 }
350 }
351 }
352 }
353 Err(e) => {
354 let err =
355 ProxyError::SecurityError(format!("JWKS endpoint returned error: {e}"));
356 error_fmt!("OidcProvider", "{}", err);
357 return Err(err);
358 }
359 },
360 Err(e) => {
361 let err =
362 ProxyError::SecurityError(format!("Failed to connect to JWKS endpoint: {e}"));
363 error_fmt!("OidcProvider", "{}", err);
364 return Err(err);
365 }
366 };
367
368 debug_fmt!(
369 "OidcProvider",
370 "JWKS refresh successful, found {} keys",
371 jwks.keys.len()
372 );
373
374 {
376 let mut w = self.jwks.write().await;
377 *w = Some(jwks);
378 }
379 {
380 let mut w = self.last_refresh.write().await;
381 *w = now;
382 }
383
384 Ok(())
385 }
386
387 pub(crate) fn jwk_to_decoding_key(&self, jwk: &Jwk) -> Result<DecodingKey, ProxyError> {
388 match &jwk.algorithm {
389 AlgorithmParameters::RSA(params) => {
390 trace_fmt!("OidcProvider", "Converting RSA JWK to decoding key");
391 DecodingKey::from_rsa_components(¶ms.n, ¶ms.e).map_err(|e| {
392 let err = ProxyError::SecurityError(format!("Invalid RSA key: {e}"));
393 error_fmt!("OidcProvider", "{}", err);
394 err
395 })
396 }
397 AlgorithmParameters::EllipticCurve(params) => {
398 trace_fmt!("OidcProvider", "Converting EC JWK to decoding key");
399 DecodingKey::from_ec_components(¶ms.x, ¶ms.y).map_err(|e| {
400 let err = ProxyError::SecurityError(format!("Invalid EC key: {e}"));
401 error_fmt!("OidcProvider", "{}", err);
402 err
403 })
404 }
405 AlgorithmParameters::OctetKey(OctetKeyParameters { value, .. }) => {
406 trace_fmt!("OidcProvider", "Converting octet JWK to decoding key");
407 Ok(DecodingKey::from_secret(value.as_bytes()))
408 }
409 AlgorithmParameters::OctetKeyPair(params) => {
410 trace_fmt!("OidcProvider", "Converting OKP JWK to decoding key");
411 DecodingKey::from_ed_components(¶ms.x).map_err(|e| {
412 let err = ProxyError::SecurityError(format!("Invalid OKP key: {e}"));
413 error_fmt!("OidcProvider", "{}", err);
414 err
415 })
416 }
417 }
418 }
419
420 pub(crate) async fn validate_token(
421 &self,
422 token: &str,
423 ) -> Result<serde_json::Value, ProxyError> {
424 let header = match decode_header(token) {
426 Ok(h) => h,
427 Err(e) => {
428 let err = ProxyError::SecurityError(format!("Invalid JWT header: {e}"));
429 warn_fmt!("OidcProvider", "{}", err);
430 return Err(err);
431 }
432 };
433
434 trace_fmt!(
435 "OidcProvider",
436 "JWT header: alg={:?}, kid={:?}",
437 header.alg,
438 header.kid
439 );
440
441 let allowed_algs = [
443 Algorithm::RS256,
444 Algorithm::RS384,
445 Algorithm::RS512,
446 Algorithm::PS256,
447 Algorithm::PS384,
448 Algorithm::PS512,
449 Algorithm::ES256,
450 Algorithm::ES384,
451 Algorithm::EdDSA,
452 Algorithm::HS256,
453 Algorithm::HS384,
454 Algorithm::HS512,
455 ];
456
457 if !allowed_algs.contains(&header.alg) {
458 let err = ProxyError::SecurityError(format!("Algorithm not allowed: {:?}", header.alg));
459 warn_fmt!("OidcProvider", "{}", err);
460 return Err(err);
461 }
462
463 match header.alg {
465 Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512 => {
466 if self.shared_secret.is_none() {
468 let err = ProxyError::SecurityError(
469 "HMAC algorithms require shared secret configuration".to_string(),
470 );
471 warn_fmt!("OidcProvider", "{}", err);
472 return Err(err);
473 }
474 if let Some(ref kid) = header.kid {
477 self.refresh_jwks().await?;
479 let jwks = self.jwks.read().await;
480 if let Some(jwks) = &*jwks {
481 if !jwks
482 .keys
483 .iter()
484 .any(|k| k.common.key_id == Some(kid.clone()))
485 {
486 let err = ProxyError::SecurityError(format!(
487 "HMAC algorithm with kid '{kid}' not found in JWKS - potential algorithm confusion attack"
488 ));
489 warn_fmt!("OidcProvider", "{}", err);
490 return Err(err);
491 }
492 }
493 }
494 }
495 _ => {
496 if header.kid.is_none() {
498 let err = ProxyError::SecurityError(
499 "Asymmetric algorithms require 'kid' (key ID) header".to_string(),
500 );
501 warn_fmt!("OidcProvider", "{}", err);
502 return Err(err);
503 }
504 }
505 }
506
507 let key = match &header.kid {
509 Some(kid) => {
510 self.refresh_jwks().await?;
512
513 let jwks = self.jwks.read().await;
515 let jwks = match &*jwks {
516 Some(j) => j,
517 None => {
518 let err = ProxyError::SecurityError("No JWKS available".to_string());
519 error_fmt!("OidcProvider", "{}", err);
520 return Err(err);
521 }
522 };
523
524 match jwks
526 .keys
527 .iter()
528 .find(|k| k.common.key_id == Some(kid.clone()))
529 {
530 Some(key) => {
531 trace_fmt!("OidcProvider", "Found key with ID {}", kid);
532 match self.jwk_to_decoding_key(key) {
533 Ok(key) => key,
534 Err(e) => {
535 error_fmt!(
536 "OidcProvider",
537 "Failed to convert JWK to decoding key: {}",
538 e
539 );
540 return Err(e);
541 }
542 }
543 }
544 None => {
545 if let Some(ref secret) = self.shared_secret {
548 if matches!(
549 header.alg,
550 Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512
551 ) {
552 trace_fmt!(
553 "OidcProvider",
554 "Using shared secret for HS* algorithm (no kid provided)"
555 );
556 DecodingKey::from_secret(secret.as_bytes())
557 } else {
558 let err = ProxyError::SecurityError(format!(
559 "Key ID {kid} not found in JWKS and algorithm {:?} requires asymmetric key",
560 header.alg
561 ));
562 warn_fmt!("OidcProvider", "{}", err);
563 return Err(err);
564 }
565 } else {
566 let err = ProxyError::SecurityError(format!(
567 "Key ID {kid} not found in JWKS"
568 ));
569 warn_fmt!("OidcProvider", "{}", err);
570 return Err(err);
571 }
572 }
573 }
574 }
575 None => {
576 if let Some(ref secret) = self.shared_secret {
579 if matches!(
580 header.alg,
581 Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512
582 ) {
583 trace_fmt!(
584 "OidcProvider",
585 "No key ID in token, using shared secret for HMAC algorithm"
586 );
587 DecodingKey::from_secret(secret.as_bytes())
588 } else {
589 let err = ProxyError::SecurityError(format!(
590 "Algorithm {:?} requires 'kid' (key ID) header for security",
591 header.alg
592 ));
593 warn_fmt!("OidcProvider", "{}", err);
594 return Err(err);
595 }
596 } else {
597 let err = ProxyError::SecurityError(
598 "No key ID in token and no shared secret configured".to_string(),
599 );
600 warn_fmt!("OidcProvider", "{}", err);
601 return Err(err);
602 }
603 }
604 };
605
606 let mut validation = Validation::new(header.alg);
608 validation.set_audience(&[&self.aud.clone().unwrap_or_default()]);
609 validation.set_issuer(&[&self.issuer]);
610
611 match decode::<serde_json::Value>(token, &key, &validation) {
613 Ok(token_data) => {
614 debug_fmt!("OidcProvider", "JWT validation successful");
615 Ok(token_data.claims)
616 }
617 Err(e) => {
618 let err = ProxyError::SecurityError(format!("JWT validation failed: {e}"));
619 warn_fmt!("OidcProvider", "{}", err);
620 Err(err)
621 }
622 }
623 }
624
625 #[allow(dead_code)]
626 pub(crate) fn validate_std_claims(&self, claims: &serde_json::Value) -> Result<(), ProxyError> {
627 if let Some(iss) = claims["iss"].as_str() {
629 if iss != self.issuer {
630 let err = ProxyError::SecurityError(format!(
631 "Invalid issuer: expected '{}', got '{}'",
632 self.issuer, iss
633 ));
634 warn_fmt!("OidcProvider", "{}", err);
635 return Err(err);
636 }
637 } else {
638 let err = ProxyError::SecurityError("Missing issuer claim".to_string());
639 warn_fmt!("OidcProvider", "{}", err);
640 return Err(err);
641 }
642
643 if let Some(ref expected_aud) = self.aud {
645 let valid_audience = match &claims["aud"] {
646 serde_json::Value::String(aud) => aud == expected_aud,
647 serde_json::Value::Array(auds) => auds
648 .iter()
649 .filter_map(|a| a.as_str())
650 .any(|a| a == expected_aud),
651 _ => false,
652 };
653
654 if !valid_audience {
655 let err = ProxyError::SecurityError(format!(
656 "Invalid audience: expected '{expected_aud}'"
657 ));
658 warn_fmt!("OidcProvider", "{}", err);
659 return Err(err);
660 }
661 }
662
663 if let Some(exp) = claims["exp"].as_i64() {
665 let now = std::time::SystemTime::now()
666 .duration_since(std::time::UNIX_EPOCH)
667 .unwrap_or_default()
668 .as_secs() as i64;
669
670 if exp <= now {
671 let err = ProxyError::SecurityError(format!(
672 "Token expired at {exp}, current time is {now}"
673 ));
674 warn_fmt!("OidcProvider", "{}", err);
675 return Err(err);
676 }
677 }
678
679 debug_fmt!("OidcProvider", "Token claims validation successful");
680 Ok(())
681 }
682
683 #[inline]
684 pub(crate) fn is_bypassed(&self, method: &str, path: &str) -> bool {
685 let bypassed = self.rules.iter().any(|r| r.matches(method, path));
686 if bypassed {
687 debug_fmt!("OidcProvider", "OIDC bypass for {} {}", method, path);
688 }
689 bypassed
690 }
691
692 pub(crate) fn extract_bearer_token<'a>(
693 &self,
694 req: &'a ProxyRequest,
695 ) -> Result<&'a str, ProxyError> {
696 debug_fmt!(
697 "OidcProvider",
698 "OIDC validating request: {} {}",
699 req.method,
700 req.path
701 );
702
703 let auth_header = if let Some(h) = req.headers.get("authorization") {
704 match h.to_str() {
705 Ok(s) => s,
706 Err(e) => {
707 let err =
708 ProxyError::SecurityError(format!("Invalid authorization header: {e}"));
709 warn_fmt!("OidcProvider", "{}", err);
710 return Err(err);
711 }
712 }
713 } else {
714 let err = ProxyError::SecurityError("Missing authorization header".to_string());
715 warn_fmt!("OidcProvider", "{}", err);
716 return Err(err);
717 };
718
719 if !auth_header.to_lowercase().starts_with(BEARER) {
720 let err = ProxyError::SecurityError(format!(
721 "Invalid authorization scheme: expected 'Bearer', got '{}'",
722 auth_header.split_whitespace().next().unwrap_or("")
723 ));
724 warn_fmt!("OidcProvider", "{}", err);
725 return Err(err);
726 }
727
728 let token = &auth_header[BEARER.len()..];
729 if token.is_empty() {
730 let err = ProxyError::SecurityError("Empty bearer token".to_string());
731 warn_fmt!("OidcProvider", "{}", err);
732 return Err(err);
733 }
734 Ok(token)
735 }
736}
737
738#[async_trait]
739impl SecurityProvider for OidcProvider {
740 fn name(&self) -> &str {
741 "OidcProvider"
742 }
743
744 fn stage(&self) -> SecurityStage {
745 SecurityStage::Pre
746 }
747
748 async fn pre(&self, req: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
749 if self.is_bypassed(&req.method.to_string(), &req.path) {
751 debug_fmt!(
752 "OidcProvider",
753 "OIDC bypass for {} {}",
754 req.method,
755 req.path
756 );
757 return Ok(req);
758 }
759
760 let token = match self.extract_bearer_token(&req) {
762 Ok(value) => value,
763 Err(value) => return Err(value),
764 };
765
766 trace_fmt!("OidcProvider", "Validating token: {}", token);
768 let claims = match self.validate_token(token).await {
769 Ok(claims) => claims,
770 Err(e) => {
771 warn_fmt!("OidcProvider", "Token validation failed: {}", e);
772 return Err(e);
773 }
774 };
775
776 {
778 let mut ctx = req.context.write().await;
779 ctx.attributes.insert(CLAIMS_ATTRIBUTE.to_string(), claims);
780 }
781
782 debug_fmt!(
783 "OidcProvider",
784 "OIDC validation successful for {} {}",
785 req.method,
786 req.path
787 );
788 Ok(req)
789 }
790}