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 pub aud: Option<String>,
70
71 #[serde(rename = "shared-secret")]
73 pub shared_secret: Option<String>,
74
75 #[serde(default)]
77 pub bypass: Vec<RouteRuleConfig>,
78}
79
80#[derive(Debug)]
82pub struct OidcProvider {
83 pub(crate) issuer: String,
85
86 pub(crate) aud: Option<String>,
88
89 pub(crate) shared_secret: Option<String>,
91
92 pub(crate) jwks_uri: String,
94
95 pub(crate) jwks: Arc<RwLock<Option<JwkSet>>>,
97
98 pub(crate) last_refresh: Arc<RwLock<tokio::time::Instant>>,
100
101 pub(crate) http: Client,
103
104 pub(crate) rules: Vec<RouteRule>,
106}
107
108impl OidcProvider {
109 pub async fn discover(cfg: OidcConfig) -> Result<Self, ProxyError> {
111 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 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 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 pub(crate) async fn refresh_jwks(&self) -> Result<(), ProxyError> {
228 let now = tokio::time::Instant::now();
229
230 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 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 {
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(¶ms.n, ¶ms.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(¶ms.x, ¶ms.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(¶ms.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 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 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 let key = match &header.kid {
370 Some(kid) => {
371 self.refresh_jwks().await?;
373
374 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 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 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 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 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 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 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 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 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 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 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 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 {
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}