1use std::{
2 sync::Arc,
3 time::{Duration, Instant},
4};
5
6use base64::Engine;
7use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header, jwk::JwkSet};
8use reqwest::Url;
9use serde_json::Value;
10use tokio::sync::RwLock;
11
12use super::types::RawOidcClaims;
13use super::{
14 OidcAuthorizationRequest, OidcClaims, OidcClientType, OidcConfig, OidcError, OidcHttpClient,
15 OidcProviderMetadata, OidcTokenExchangeRequest, OidcUserInfo, PkcePair, ReqwestOidcHttpClient,
16};
17
18const JWKS_REFRESH_COOLDOWN: Duration = Duration::from_secs(30);
19
20#[derive(Debug, Default)]
21struct OidcCache {
22 metadata: Option<OidcProviderMetadata>,
23 jwks: Option<Arc<JwkSet>>,
24 last_forced_jwks_refresh: Option<Instant>,
25}
26
27#[derive(Debug)]
29pub struct OidcClient<H = ReqwestOidcHttpClient> {
30 config: OidcConfig,
31 http_client: H,
32 cache: Arc<RwLock<OidcCache>>,
33}
34
35impl OidcClient<ReqwestOidcHttpClient> {
36 pub fn new(config: OidcConfig) -> Result<Self, OidcError> {
41 Self::with_http_client(config, ReqwestOidcHttpClient::new()?)
42 }
43}
44
45impl<H> OidcClient<H>
46where
47 H: OidcHttpClient,
48{
49 pub fn with_http_client(config: OidcConfig, http_client: H) -> Result<Self, OidcError> {
54 config.validate()?;
55 Ok(Self {
56 config,
57 http_client,
58 cache: Arc::new(RwLock::new(OidcCache::default())),
59 })
60 }
61
62 pub async fn discover(&self) -> Result<OidcProviderMetadata, OidcError> {
67 let cached_metadata = self.cache.read().await.metadata.clone();
68 if let Some(metadata) = cached_metadata {
69 return Ok(metadata);
70 }
71
72 let issuer = self.config.issuer.trim_end_matches('/');
73 let url = format!("{issuer}/.well-known/openid-configuration");
74 let json = self.http_client.get_json(&url, None).await?;
75 let metadata = serde_json::from_value::<OidcProviderMetadata>(json).map_err(|error| {
76 OidcError::Discovery(format!("invalid discovery document: {error}"))
77 })?;
78 if metadata.issuer.trim_end_matches('/') != self.config.issuer.trim_end_matches('/') {
79 return Err(OidcError::Discovery(
80 "provider issuer did not exactly match configured issuer".into(),
81 ));
82 }
83 if !metadata
84 .response_types_supported
85 .iter()
86 .any(|value| value == "code")
87 {
88 return Err(OidcError::Discovery(
89 "provider must support the authorization code flow".into(),
90 ));
91 }
92 if matches!(self.config.client_type, OidcClientType::Public)
93 && !metadata
94 .code_challenge_methods_supported
95 .iter()
96 .any(|method| method == "S256")
97 {
98 return Err(OidcError::Discovery(
99 "public clients require PKCE S256 support".into(),
100 ));
101 }
102
103 self.cache.write().await.metadata = Some(metadata.clone());
104 Ok(metadata)
105 }
106
107 async fn jwks(&self, force_refresh: bool) -> Result<Arc<JwkSet>, OidcError> {
108 if !force_refresh && let Some(jwks) = self.cache.read().await.jwks.clone() {
109 return Ok(jwks);
110 }
111
112 if force_refresh {
113 let mut cache = self.cache.write().await;
114 if let Some(jwks) = cache.jwks.clone()
115 && cache
116 .last_forced_jwks_refresh
117 .is_some_and(|last_refresh| last_refresh.elapsed() < JWKS_REFRESH_COOLDOWN)
118 {
119 return Ok(jwks);
120 }
121 cache.last_forced_jwks_refresh = Some(Instant::now());
122 }
123
124 let metadata = self.discover().await?;
125 let json = self.http_client.get_json(&metadata.jwks_uri, None).await?;
126 let jwks = serde_json::from_value::<JwkSet>(json)
127 .map_err(|error| OidcError::Discovery(format!("invalid JWKS document: {error}")))?;
128 let jwks = Arc::new(jwks);
129 self.cache.write().await.jwks = Some(Arc::clone(&jwks));
130 Ok(jwks)
131 }
132
133 pub async fn build_authorization_request(
138 &self,
139 scopes: &[&str],
140 ) -> Result<OidcAuthorizationRequest, OidcError> {
141 let metadata = self.discover().await?;
142 let state = random_urlsafe(24);
143 let nonce = random_urlsafe(24);
144 let pkce = Some(PkcePair::generate());
145
146 let mut url = Url::parse(&metadata.authorization_endpoint).map_err(|error| {
147 OidcError::Discovery(format!("invalid authorization endpoint: {error}"))
148 })?;
149 {
150 let mut query = url.query_pairs_mut();
151 query.append_pair("response_type", "code");
152 query.append_pair("client_id", &self.config.client_id);
153 query.append_pair("redirect_uri", &self.config.redirect_uri);
154 query.append_pair("scope", &scopes.join(" "));
155 query.append_pair("state", &state);
156 query.append_pair("nonce", &nonce);
157 if let Some(pkce) = &pkce {
158 query.append_pair("code_challenge", &pkce.challenge);
159 query.append_pair("code_challenge_method", pkce.method);
160 }
161 }
162
163 Ok(OidcAuthorizationRequest {
164 url: url.to_string(),
165 state,
166 nonce,
167 pkce,
168 })
169 }
170
171 pub async fn build_token_exchange_request(
176 &self,
177 pending: &OidcAuthorizationRequest,
178 code: &str,
179 returned_state: &str,
180 code_verifier: Option<&str>,
181 ) -> Result<OidcTokenExchangeRequest, OidcError> {
182 if pending.state != returned_state {
183 return Err(OidcError::StateMismatch);
184 }
185
186 let verifier = code_verifier
187 .map(ToOwned::to_owned)
188 .or_else(|| pending.pkce.as_ref().map(|pkce| pkce.verifier.clone()));
189
190 if matches!(self.config.client_type, OidcClientType::Public) && verifier.is_none() {
191 return Err(OidcError::MissingPkce);
192 }
193
194 let metadata = self.discover().await?;
195 Ok(OidcTokenExchangeRequest {
196 token_endpoint: metadata.token_endpoint,
197 code: code.to_string(),
198 redirect_uri: self.config.redirect_uri.clone(),
199 state: returned_state.to_string(),
200 code_verifier: verifier,
201 })
202 }
203
204 pub async fn validate_id_token(
209 &self,
210 id_token: &str,
211 expected_nonce: Option<&str>,
212 ) -> Result<OidcClaims, OidcError> {
213 let metadata = self.discover().await?;
214 let header = decode_header(id_token)
215 .map_err(|error| OidcError::InvalidToken(format!("invalid token header: {error}")))?;
216
217 if !self.config.allowed_algorithms.contains(&header.alg) {
218 return Err(OidcError::UnsupportedAlgorithm(format!("{:?}", header.alg)));
219 }
220 let alg_name = format!("{:?}", header.alg);
221 if !metadata.id_token_signing_alg_values_supported.is_empty()
222 && !metadata
223 .id_token_signing_alg_values_supported
224 .iter()
225 .any(|value| value == &alg_name)
226 {
227 return Err(OidcError::UnsupportedAlgorithm(alg_name));
228 }
229
230 let jwk = self.select_jwk(header.kid.as_deref()).await?;
231 let decoding_key = DecodingKey::from_jwk(&jwk).map_err(|error| {
232 OidcError::InvalidToken(format!("could not build decoding key from JWK: {error}"))
233 })?;
234 let claims = decode::<Value>(
235 id_token,
236 &decoding_key,
237 &oidc_validation(&self.config, header.alg),
238 )
239 .map_err(|error| map_oidc_jwt_error(&error))?
240 .claims;
241
242 let raw_claims = serde_json::from_value::<RawOidcClaims>(claims)
243 .map_err(|error| OidcError::InvalidToken(format!("invalid OIDC claims: {error}")))?;
244 let iat = raw_claims
245 .iat
246 .ok_or_else(|| OidcError::MissingClaim("iat".into()))?;
247 if let Some(expected_nonce) = expected_nonce
248 && raw_claims.nonce.as_deref() != Some(expected_nonce)
249 {
250 return Err(OidcError::NonceMismatch);
251 }
252
253 Ok(OidcClaims {
254 sub: raw_claims.sub,
255 iss: raw_claims.iss,
256 aud: raw_claims.aud.into_vec(),
257 exp: raw_claims.exp,
258 iat,
259 nbf: raw_claims.nbf,
260 nonce: raw_claims.nonce,
261 email: raw_claims.email,
262 email_verified: raw_claims.email_verified,
263 name: raw_claims.name,
264 })
265 }
266
267 async fn select_jwk(&self, kid: Option<&str>) -> Result<jsonwebtoken::jwk::Jwk, OidcError> {
268 let kid = kid.ok_or_else(|| {
269 OidcError::InvalidToken("token header is missing required kid".into())
270 })?;
271 for force_refresh in [false, true] {
272 let jwks = self.jwks(force_refresh).await?;
273 if let Some(jwk) = jwks
274 .keys
275 .iter()
276 .find(|jwk| jwk.common.key_id.as_deref() == Some(kid))
277 .cloned()
278 {
279 return Ok(jwk);
280 }
281 }
282 Err(OidcError::InvalidToken(
283 "no matching JWK found for token header".into(),
284 ))
285 }
286
287 pub async fn fetch_userinfo(&self, access_token: &str) -> Result<OidcUserInfo, OidcError> {
292 let metadata = self.discover().await?;
293 let endpoint = metadata.userinfo_endpoint.ok_or_else(|| {
294 OidcError::Discovery("provider does not expose a userinfo endpoint".into())
295 })?;
296 let json = self
297 .http_client
298 .get_json(&endpoint, Some(access_token))
299 .await?;
300 serde_json::from_value(json).map_err(|error| {
301 OidcError::ProviderUnreachable(format!("invalid userinfo response: {error}"))
302 })
303 }
304}
305
306fn oidc_validation(config: &OidcConfig, algorithm: Algorithm) -> Validation {
307 let mut validation = Validation::new(algorithm);
308 validation.algorithms = vec![algorithm];
309 validation.set_issuer(&[config.issuer.as_str()]);
310 validation.set_audience(&config.audience);
311 validation.set_required_spec_claims(&["exp", "iss", "aud", "sub"]);
312 validation.validate_nbf = true;
315 validation.leeway = config.clock_skew.as_secs();
316 validation
317}
318
319fn map_oidc_jwt_error(error: &jsonwebtoken::errors::Error) -> OidcError {
320 match error.kind() {
321 jsonwebtoken::errors::ErrorKind::ExpiredSignature => {
322 OidcError::InvalidToken("OIDC ID token has expired".into())
323 }
324 jsonwebtoken::errors::ErrorKind::MissingRequiredClaim(claim) => {
325 OidcError::MissingClaim(claim.clone())
326 }
327 _ => OidcError::InvalidToken(error.to_string()),
328 }
329}
330
331fn random_urlsafe(len: usize) -> String {
332 let mut bytes = vec![0_u8; len];
333 rand::fill(bytes.as_mut_slice());
334 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
335}
336
337pub async fn validate_id_token(
342 config: &OidcConfig,
343 id_token: &str,
344) -> Result<OidcClaims, OidcError> {
345 OidcClient::new(config.clone())?
346 .validate_id_token(id_token, None)
347 .await
348}
349
350#[cfg(test)]
351mod tests {
352 use std::{
353 collections::HashMap,
354 sync::{
355 Arc,
356 atomic::{AtomicUsize, Ordering},
357 },
358 };
359
360 use async_trait::async_trait;
361 use jsonwebtoken::{EncodingKey, Header, encode};
362
363 use super::*;
364
365 const ISSUER: &str = "https://issuer.example";
366 const CLIENT_ID: &str = "client-123";
367 const REDIRECT_URI: &str = "https://app.example/callback";
368 const RSA_PRIVATE_KEY: &str = include_str!(concat!(
369 env!("CARGO_MANIFEST_DIR"),
370 "/testdata/rsa_private_key.pem"
371 ));
372 const JWKS_JSON: &str = r#"{
373 "keys": [{
374 "kty": "RSA",
375 "kid": "rsa-1",
376 "use": "sig",
377 "alg": "RS256",
378 "n": "oT6vqY7IFYxu8LYnCgPsGZICOgRF57XQI9wAoILWbGwIBTzZMi_KuWWtYCo-Ph3LhN2vsIUQkKwT1-kob3IfBoCh9kN0iO6f_XLXgAOCCtVyrt5bjLyFJvtYcFecfh80LvdEeL8VE7Fxnd6CoRv1zNszWFxsfSCeWPyWQefJjLcWOqRg5zyFuP5yHwGwzEI3wLKhPFC2ufTupu4GRcFpjKpM3ZxoWw2BEaPcSmKMFWBGd7lsKaBe6TfLvL3pDwzBHSngURGInXrTQGmw-KRpDrcv5vlUD7QwVWP3JWGW66RtWH5qS5MP2eeOOSd48_Yccbx1kbGZ-NQtk6fSOTwd_Q",
379 "e": "AQAB"
380 }]
381 }"#;
382
383 fn random_nonce() -> String {
385 format!("nonce-{:016x}", rand::random::<u64>())
386 }
387
388 #[derive(Debug, Clone)]
389 struct MockHttpClient {
390 responses: Arc<HashMap<String, Value>>,
391 request_counts: Arc<HashMap<String, AtomicUsize>>,
392 }
393
394 #[async_trait]
395 impl OidcHttpClient for MockHttpClient {
396 async fn get_json(
397 &self,
398 url: &str,
399 _bearer_token: Option<&str>,
400 ) -> Result<Value, OidcError> {
401 if let Some(counter) = self.request_counts.get(url) {
402 counter.fetch_add(1, Ordering::Relaxed);
403 }
404 self.responses.get(url).cloned().ok_or_else(|| {
405 OidcError::ProviderUnreachable(format!("no response configured for {url}"))
406 })
407 }
408 }
409
410 fn mock_client() -> OidcClient<MockHttpClient> {
411 let responses = HashMap::from([
412 (
413 format!("{ISSUER}/.well-known/openid-configuration"),
414 serde_json::json!({
415 "issuer": ISSUER,
416 "authorization_endpoint": format!("{ISSUER}/authorize"),
417 "token_endpoint": format!("{ISSUER}/token"),
418 "jwks_uri": format!("{ISSUER}/jwks"),
419 "userinfo_endpoint": format!("{ISSUER}/userinfo"),
420 "response_types_supported": ["code"],
421 "code_challenge_methods_supported": ["S256"],
422 "id_token_signing_alg_values_supported": ["RS256"]
423 }),
424 ),
425 (
426 format!("{ISSUER}/jwks"),
427 serde_json::from_str(JWKS_JSON).unwrap(),
428 ),
429 (
430 format!("{ISSUER}/userinfo"),
431 serde_json::json!({
432 "sub": "user-123",
433 "email": "user@example.com",
434 "email_verified": true,
435 "name": "Example User"
436 }),
437 ),
438 ]);
439
440 OidcClient::with_http_client(
441 OidcConfig::new(ISSUER, CLIENT_ID, REDIRECT_URI, OidcClientType::Public),
442 MockHttpClient {
443 responses: Arc::new(responses),
444 request_counts: Arc::new(HashMap::from([
445 (
446 format!("{ISSUER}/.well-known/openid-configuration"),
447 AtomicUsize::new(0),
448 ),
449 (format!("{ISSUER}/jwks"), AtomicUsize::new(0)),
450 (format!("{ISSUER}/userinfo"), AtomicUsize::new(0)),
451 ])),
452 },
453 )
454 .unwrap()
455 }
456
457 fn mock_client_with_discovery(discovery: Value) -> OidcClient<MockHttpClient> {
458 let responses = HashMap::from([(
459 format!("{ISSUER}/.well-known/openid-configuration"),
460 discovery,
461 )]);
462
463 OidcClient::with_http_client(
464 OidcConfig::new(ISSUER, CLIENT_ID, REDIRECT_URI, OidcClientType::Public),
465 MockHttpClient {
466 responses: Arc::new(responses),
467 request_counts: Arc::new(HashMap::new()),
468 },
469 )
470 .unwrap()
471 }
472
473 fn issue_token(nonce: &str) -> String {
474 let now = std::time::SystemTime::now()
475 .duration_since(std::time::UNIX_EPOCH)
476 .unwrap()
477 .as_secs();
478 let mut header = Header::new(Algorithm::RS256);
479 header.kid = Some("rsa-1".into());
480 encode(
481 &header,
482 &serde_json::json!({
483 "sub": "user-123",
484 "iss": ISSUER,
485 "aud": [CLIENT_ID],
486 "exp": now + 3600,
487 "nbf": now.saturating_sub(1),
488 "iat": now,
489 "nonce": nonce,
490 "email": "user@example.com",
491 "email_verified": true,
492 "name": "Example User"
493 }),
494 &EncodingKey::from_rsa_pem(RSA_PRIVATE_KEY.as_bytes()).unwrap(),
495 )
496 .unwrap()
497 }
498
499 fn issue_token_without_nbf(nonce: &str) -> String {
500 let now = std::time::SystemTime::now()
501 .duration_since(std::time::UNIX_EPOCH)
502 .unwrap()
503 .as_secs();
504 let mut header = Header::new(Algorithm::RS256);
505 header.kid = Some("rsa-1".into());
506 encode(
507 &header,
508 &serde_json::json!({
509 "sub": "user-123",
510 "iss": ISSUER,
511 "aud": [CLIENT_ID],
512 "exp": now + 3600,
513 "iat": now,
514 "nonce": nonce,
515 "email": "user@example.com",
516 "email_verified": true,
517 "name": "Example User"
518 }),
519 &EncodingKey::from_rsa_pem(RSA_PRIVATE_KEY.as_bytes()).unwrap(),
520 )
521 .unwrap()
522 }
523
524 #[tokio::test]
525 async fn authorization_request_includes_pkce_state_and_nonce() {
526 let client = mock_client();
527 let request = client
528 .build_authorization_request(&["openid", "profile", "email"])
529 .await
530 .unwrap();
531
532 assert!(request.url.contains("response_type=code"));
533 assert!(request.url.contains("code_challenge="));
534 assert!(!request.state.is_empty());
535 assert!(!request.nonce.is_empty());
536 }
537
538 #[tokio::test]
539 async fn discovery_rejects_issuer_mismatch_and_missing_authorization_code_support() {
540 let issuer_mismatch = mock_client_with_discovery(serde_json::json!({
541 "issuer": "https://other-issuer.example",
542 "authorization_endpoint": format!("{ISSUER}/authorize"),
543 "token_endpoint": format!("{ISSUER}/token"),
544 "jwks_uri": format!("{ISSUER}/jwks"),
545 "response_types_supported": ["code"],
546 "code_challenge_methods_supported": ["S256"],
547 "id_token_signing_alg_values_supported": ["RS256"]
548 }));
549 assert!(matches!(
550 issuer_mismatch.discover().await,
551 Err(OidcError::Discovery(message))
552 if message.contains("issuer did not exactly match")
553 ));
554
555 let no_code_flow = mock_client_with_discovery(serde_json::json!({
556 "issuer": ISSUER,
557 "authorization_endpoint": format!("{ISSUER}/authorize"),
558 "token_endpoint": format!("{ISSUER}/token"),
559 "jwks_uri": format!("{ISSUER}/jwks"),
560 "response_types_supported": ["token"],
561 "code_challenge_methods_supported": ["S256"],
562 "id_token_signing_alg_values_supported": ["RS256"]
563 }));
564 assert!(matches!(
565 no_code_flow.discover().await,
566 Err(OidcError::Discovery(message))
567 if message.contains("authorization code flow")
568 ));
569 }
570
571 #[tokio::test]
572 async fn public_client_discovery_requires_s256_pkce_support() {
573 let client = mock_client_with_discovery(serde_json::json!({
574 "issuer": ISSUER,
575 "authorization_endpoint": format!("{ISSUER}/authorize"),
576 "token_endpoint": format!("{ISSUER}/token"),
577 "jwks_uri": format!("{ISSUER}/jwks"),
578 "response_types_supported": ["code"],
579 "code_challenge_methods_supported": ["plain"],
580 "id_token_signing_alg_values_supported": ["RS256"]
581 }));
582
583 assert!(matches!(
584 client.discover().await,
585 Err(OidcError::Discovery(message)) if message.contains("PKCE S256")
586 ));
587 }
588
589 #[tokio::test]
590 async fn state_mismatch_is_rejected() {
591 let client = mock_client();
592 let request = client
593 .build_authorization_request(&["openid"])
594 .await
595 .unwrap();
596 let result = client
597 .build_token_exchange_request(&request, "code-123", "wrong-state", None)
598 .await;
599 assert_eq!(result.unwrap_err(), OidcError::StateMismatch);
600 }
601
602 #[tokio::test]
603 async fn pkce_missing_is_rejected_for_public_client() {
604 let client = mock_client();
605 let request = OidcAuthorizationRequest {
606 url: "https://issuer.example/authorize".into(),
607 state: "state-123".into(),
608 nonce: random_nonce(),
609 pkce: None,
610 };
611 let result = client
612 .build_token_exchange_request(&request, "code-123", "state-123", None)
613 .await;
614 assert_eq!(result.unwrap_err(), OidcError::MissingPkce);
615 }
616
617 #[tokio::test]
618 async fn token_exchange_uses_pending_pkce_when_state_matches() {
619 let client = mock_client();
620 let pending = client
621 .build_authorization_request(&["openid"])
622 .await
623 .unwrap();
624 let verifier = pending.pkce.as_ref().unwrap().verifier.clone();
625
626 let exchange = client
627 .build_token_exchange_request(&pending, "code-123", &pending.state, None)
628 .await
629 .unwrap();
630
631 assert_eq!(exchange.token_endpoint, format!("{ISSUER}/token"));
632 assert_eq!(exchange.code, "code-123");
633 assert_eq!(exchange.redirect_uri, REDIRECT_URI);
634 assert_eq!(exchange.state, pending.state);
635 assert_eq!(exchange.code_verifier.as_deref(), Some(verifier.as_str()));
636 }
637
638 #[tokio::test]
639 async fn validate_id_token_rejects_malformed_token_before_fetching_jwks() {
640 let client = mock_client();
641
642 let error = client
643 .validate_id_token("not-a-jwt", Some(&random_nonce()))
644 .await
645 .unwrap_err();
646
647 assert!(matches!(
648 error,
649 OidcError::InvalidToken(message) if message.contains("invalid token header")
650 ));
651 }
652
653 #[tokio::test]
654 async fn nonce_mismatch_is_rejected() {
655 let client = mock_client();
656 let expected = random_nonce();
657 let wrong = format!("{expected}-mismatch");
658 let token = issue_token(&expected);
659 let result = client.validate_id_token(&token, Some(&wrong)).await;
660 assert_eq!(result.unwrap_err(), OidcError::NonceMismatch);
661 }
662
663 #[tokio::test]
664 async fn valid_id_token_and_userinfo_roundtrip() {
665 let client = mock_client();
666 let nonce = random_nonce();
667 let token = issue_token(&nonce);
668 let claims = client
669 .validate_id_token(&token, Some(&nonce))
670 .await
671 .unwrap();
672 let userinfo = client.fetch_userinfo("opaque-access-token").await.unwrap();
673
674 assert_eq!(claims.sub, "user-123");
675 assert_eq!(claims.email.as_deref(), Some("user@example.com"));
676 assert_eq!(userinfo.name.as_deref(), Some("Example User"));
677 }
678
679 #[tokio::test]
680 async fn valid_id_token_without_nbf_is_accepted() {
681 let client = mock_client();
682 let nonce = random_nonce();
683 let token = issue_token_without_nbf(&nonce);
684
685 let claims = client
686 .validate_id_token(&token, Some(&nonce))
687 .await
688 .unwrap();
689
690 assert_eq!(claims.sub, "user-123");
691 assert!(claims.nbf.is_none());
692 }
693
694 #[tokio::test]
695 async fn id_token_without_kid_is_rejected() {
696 let client = mock_client();
697 let now = std::time::SystemTime::now()
698 .duration_since(std::time::UNIX_EPOCH)
699 .unwrap()
700 .as_secs();
701 let nonce = random_nonce();
702 let token = encode(
703 &Header::new(Algorithm::RS256),
704 &serde_json::json!({
705 "sub": "user-123",
706 "iss": ISSUER,
707 "aud": [CLIENT_ID],
708 "exp": now + 3600,
709 "nbf": now.saturating_sub(1),
710 "iat": now,
711 "nonce": nonce
712 }),
713 &EncodingKey::from_rsa_pem(RSA_PRIVATE_KEY.as_bytes()).unwrap(),
714 )
715 .unwrap();
716
717 let result = client.validate_id_token(&token, Some(&nonce)).await;
718 assert_eq!(
719 result.unwrap_err(),
720 OidcError::InvalidToken("token header is missing required kid".into())
721 );
722 }
723
724 #[tokio::test]
725 async fn fetch_userinfo_fails_closed_when_endpoint_is_absent_or_malformed() {
726 let without_userinfo = mock_client_with_discovery(serde_json::json!({
727 "issuer": ISSUER,
728 "authorization_endpoint": format!("{ISSUER}/authorize"),
729 "token_endpoint": format!("{ISSUER}/token"),
730 "jwks_uri": format!("{ISSUER}/jwks"),
731 "response_types_supported": ["code"],
732 "code_challenge_methods_supported": ["S256"],
733 "id_token_signing_alg_values_supported": ["RS256"]
734 }));
735 assert!(matches!(
736 without_userinfo.fetch_userinfo("token").await,
737 Err(OidcError::Discovery(message))
738 if message.contains("userinfo endpoint")
739 ));
740
741 let mut responses = HashMap::new();
742 responses.insert(
743 format!("{ISSUER}/.well-known/openid-configuration"),
744 serde_json::json!({
745 "issuer": ISSUER,
746 "authorization_endpoint": format!("{ISSUER}/authorize"),
747 "token_endpoint": format!("{ISSUER}/token"),
748 "jwks_uri": format!("{ISSUER}/jwks"),
749 "userinfo_endpoint": format!("{ISSUER}/userinfo"),
750 "response_types_supported": ["code"],
751 "code_challenge_methods_supported": ["S256"],
752 "id_token_signing_alg_values_supported": ["RS256"]
753 }),
754 );
755 responses.insert(
756 format!("{ISSUER}/userinfo"),
757 serde_json::json!({"email": "missing-sub@example.com"}),
758 );
759 let malformed_userinfo = OidcClient::with_http_client(
760 OidcConfig::new(ISSUER, CLIENT_ID, REDIRECT_URI, OidcClientType::Public),
761 MockHttpClient {
762 responses: Arc::new(responses),
763 request_counts: Arc::new(HashMap::new()),
764 },
765 )
766 .unwrap();
767
768 assert!(matches!(
769 malformed_userinfo.fetch_userinfo("token").await,
770 Err(OidcError::ProviderUnreachable(message))
771 if message.contains("invalid userinfo response")
772 ));
773 }
774
775 #[tokio::test]
776 async fn jwks_forced_refresh_is_throttled_after_unknown_kid() {
777 let client = mock_client();
778 let now = std::time::SystemTime::now()
779 .duration_since(std::time::UNIX_EPOCH)
780 .unwrap()
781 .as_secs();
782 let mut header = Header::new(Algorithm::RS256);
783 header.kid = Some("missing-kid".into());
784 let nonce = random_nonce();
785 let token = encode(
786 &header,
787 &serde_json::json!({
788 "sub": "user-123",
789 "iss": ISSUER,
790 "aud": [CLIENT_ID],
791 "exp": now + 3600,
792 "nbf": now.saturating_sub(1),
793 "iat": now,
794 "nonce": nonce
795 }),
796 &EncodingKey::from_rsa_pem(RSA_PRIVATE_KEY.as_bytes()).unwrap(),
797 )
798 .unwrap();
799
800 let jwks_url = format!("{ISSUER}/jwks");
801 let counter = client
802 .http_client
803 .request_counts
804 .get(&jwks_url)
805 .expect("jwks counter must exist");
806
807 let first = client.validate_id_token(&token, Some(&nonce)).await;
808 assert_eq!(
809 first.unwrap_err(),
810 OidcError::InvalidToken("no matching JWK found for token header".into())
811 );
812 assert_eq!(counter.load(Ordering::Relaxed), 2);
813
814 let second = client.validate_id_token(&token, Some(&nonce)).await;
815 assert_eq!(
816 second.unwrap_err(),
817 OidcError::InvalidToken("no matching JWK found for token header".into())
818 );
819 assert_eq!(counter.load(Ordering::Relaxed), 2);
820 }
821}