1use 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}, security::{SecurityProvider, SecurityStage}};
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 log::trace!("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#[derive(Debug, Clone, Deserialize)]
53pub struct OidcConfig {
54 #[serde(rename = "issuer-uri")]
55 pub issuer_uri: String,
56
57 pub aud: Option<String>,
59
60 #[serde(rename = "shared-secret")]
62 pub shared_secret: Option<String>,
63
64 #[serde(default)]
66 pub bypass: Vec<RouteRuleConfig>,
67}
68
69#[derive(Debug)]
71pub struct OidcProvider {
72 issuer: String,
74
75 aud: Option<String>,
77
78 shared_secret: Option<String>,
80
81 jwks_uri: String,
83
84 jwks: Arc<RwLock<Option<JwkSet>>>,
86
87 last_refresh: Arc<RwLock<tokio::time::Instant>>,
89
90 http: Client,
92
93 rules: Vec<RouteRule>,
95}
96
97impl OidcProvider {
98 pub async fn discover(cfg: OidcConfig) -> Result<Self, ProxyError> {
100 log::debug!("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 log::error!("{}", 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 log::error!("{}", 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 log::error!("{}", 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 log::error!("{}", err);
144 return Err(err);
145 }
146 };
147
148 log::debug!("OIDC discovery successful, JWKS URI: {}", meta.jwks_uri);
149
150 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 log::error!("{}", err);
166 return Err(err);
167 }
168 },
169 });
170 log::debug!("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 log::error!("{}", 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 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 log::trace!("JWKS cache still fresh, skipping refresh");
205 return Ok(());
206 }
207
208 log::debug!("Refreshing JWKS from {}", self.jwks_uri);
209
210 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 log::error!("{}", 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 log::error!("{}", 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 log::error!("{}", err);
240 return Err(err);
241 }
242 };
243
244 log::debug!("JWKS refresh successful, found {} keys", jwks.keys.len());
245
246 {
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, _alg: Algorithm) -> Result<DecodingKey, ProxyError> {
260 match &jwk.algorithm {
261 AlgorithmParameters::RSA(params) => {
262 log::trace!("Converting RSA JWK to decoding key");
263 DecodingKey::from_rsa_components(¶ms.n, ¶ms.e)
264 .map_err(|e| {
265 let err = ProxyError::SecurityError(format!("Invalid RSA key: {}", e));
266 log::error!("{}", err);
267 err
268 })
269 }
270 AlgorithmParameters::EllipticCurve(params) => {
271 log::trace!("Converting EC JWK to decoding key");
272 DecodingKey::from_ec_components(¶ms.x, ¶ms.y)
273 .map_err(|e| {
274 let err = ProxyError::SecurityError(format!("Invalid EC key: {}", e));
275 log::error!("{}", err);
276 err
277 })
278 }
279 AlgorithmParameters::OctetKey(OctetKeyParameters { value, .. }) => {
280 log::trace!("Converting octet JWK to decoding key");
281 Ok(DecodingKey::from_secret(value.as_bytes()))
282 }
283 AlgorithmParameters::OctetKeyPair(params) => {
284 log::trace!("Converting OKP JWK to decoding key");
285 DecodingKey::from_ed_components(¶ms.x)
286 .map_err(|e| {
287 let err = ProxyError::SecurityError(format!("Invalid OKP key: {}", e));
288 log::error!("{}", err);
289 err
290 })
291 }
292 }
293 }
294
295 async fn validate_token(&self, token: &str) -> Result<serde_json::Value, ProxyError> {
296 let header = match decode_header(token) {
298 Ok(h) => h,
299 Err(e) => {
300 let err = ProxyError::SecurityError(format!("Invalid JWT header: {}", e));
301 log::warn!("{}", err);
302 return Err(err);
303 }
304 };
305
306 log::trace!("JWT header: alg={:?}, kid={:?}", header.alg, header.kid);
307
308 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 log::warn!("{}", err);
322 return Err(err);
323 }
324
325 self.refresh_jwks().await?;
327
328 let key = match &header.kid {
330 Some(kid) => {
331 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 log::error!("{}", err);
338 return Err(err);
339 }
340 };
341
342 match jwks.keys.iter().find(|k| k.common.key_id == Some(kid.clone())) {
344 Some(key) => {
345 log::trace!("Found key with ID {}", kid);
346 match self.jwk_to_decoding_key(key, header.alg) {
347 Ok(key) => key,
348 Err(e) => {
349 log::error!("Failed to convert JWK to decoding key: {}", e);
350 return Err(e);
351 }
352 }
353 }
354 None => {
355 if let Some(ref secret) = self.shared_secret {
357 if matches!(header.alg, Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512) {
358 log::trace!("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 log::warn!("{}", err);
363 return Err(err);
364 }
365 } else {
366 let err = ProxyError::SecurityError(format!("Key ID {} not found in JWKS", kid));
367 log::warn!("{}", err);
368 return Err(err);
369 }
370 }
371 }
372 }
373 None => {
374 if let Some(ref secret) = self.shared_secret {
376 log::trace!("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 log::warn!("{}", err);
381 return Err(err);
382 }
383 }
384 };
385
386 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 match decode::<serde_json::Value>(token, &key, &validation) {
393 Ok(token_data) => {
394 log::debug!("JWT validation successful");
395 Ok(token_data.claims)
396 }
397 Err(e) => {
398 let err = ProxyError::SecurityError(format!("JWT validation failed: {}", e));
399 log::warn!("{}", err);
400 Err(err)
401 }
402 }
403 }
404
405 fn validate_std_claims(&self, claims: &serde_json::Value) -> Result<(), ProxyError> {
406 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 log::warn!("{}", err);
413 return Err(err);
414 }
415 } else {
416 let err = ProxyError::SecurityError("Missing issuer claim".to_string());
417 log::warn!("{}", err);
418 return Err(err);
419 }
420
421 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 log::warn!("{}", err);
436 return Err(err);
437 }
438 }
439
440 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 log::warn!("{}", err);
452 return Err(err);
453 }
454 }
455
456 log::debug!("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 log::debug!("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 if self.is_bypassed(&req.method.to_string(), &req.path) {
479 log::debug!("OIDC bypass for {} {}", req.method, req.path);
480 return Ok(req);
481 }
482
483 log::debug!("OIDC validating request: {} {}", req.method, req.path);
484
485 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 log::warn!("{}", err);
494 return Err(err);
495 }
496 },
497 None => {
498 let err = ProxyError::SecurityError("Missing authorization header".to_string());
499 log::warn!("{}", 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 log::warn!("{}", 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 log::warn!("{}", err);
517 return Err(err);
518 }
519
520 log::trace!("Validating token: {}", token);
522 let claims = match self.validate_token(token).await {
523 Ok(claims) => claims,
524 Err(e) => {
525 log::warn!("Token validation failed: {}", e);
526 return Err(e);
527 }
528 };
529
530 {
532 let mut ctx = req.context.write().await;
533 ctx.attributes.insert(CLAIMS_ATTRIBUTE.to_string(), claims);
534 }
535
536 log::debug!("OIDC validation successful for {} {}", req.method, req.path);
537 Ok(req)
538 }
539}