1use std::collections::HashMap;
2use std::sync::Arc;
3use std::time::Duration;
4use std::time::Instant;
5
6use axum::body::Body;
7use axum::extract::State;
8use axum::http::{HeaderMap, Request, StatusCode, header};
9use axum::middleware::Next;
10use axum::response::Response;
11use jsonwebtoken::jwk::{Jwk, JwkSet, KeyAlgorithm, KeyOperations, PublicKeyUse};
12use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
13use serde::Deserialize;
14use serde_json::json;
15use tokio::sync::RwLock;
16use tracing::{debug, warn};
17use url::{Host, Url};
18
19use agentic_core::utils::common::uuid7_str;
20
21const OIDC_HTTP_TIMEOUT: Duration = Duration::from_secs(10);
22const JWKS_REFRESH_COOLDOWN: Duration = Duration::from_secs(30);
23const JWKS_REFRESH_COALESCE_WINDOW: Duration = Duration::from_secs(1);
24const JWKS_REFRESH_WAIT_TIMEOUT: Duration = Duration::from_secs(1);
25const DEFAULT_JWKS_TTL: Duration = Duration::from_secs(300);
26const MAX_JWKS_TTL: Duration = Duration::from_secs(3600);
27const MAX_PROVIDER_RESPONSE_BYTES: usize = 1024 * 1024;
28const MAX_JWKS_KEYS: usize = 100;
29const JWT_CLOCK_SKEW_SECONDS: u64 = 60;
30const SUPPORTED_VERIFICATION_ALGORITHMS: [Algorithm; 9] = [
31 Algorithm::ES256,
32 Algorithm::ES384,
33 Algorithm::RS256,
34 Algorithm::RS384,
35 Algorithm::RS512,
36 Algorithm::PS256,
37 Algorithm::PS384,
38 Algorithm::PS512,
39 Algorithm::EdDSA,
40];
41
42pub(crate) const ANTHROPIC_MESSAGES_PATH: &str = "/v1/messages";
43pub(crate) const ANTHROPIC_COUNT_TOKENS_PATH: &str = "/v1/messages/count_tokens";
44
45#[derive(Clone)]
46pub struct OidcConfig {
47 issuer: Url,
48 issuer_value: String,
49 audience: String,
50 allows_loopback_http: bool,
51}
52
53impl OidcConfig {
54 pub fn new(issuer: &str, audience: &str) -> Result<Self, OidcAuthError> {
61 let issuer_value = issuer.trim().to_owned();
62 let issuer = Url::parse(&issuer_value).map_err(OidcAuthError::InvalidIssuer)?;
63 let allows_loopback_http = is_loopback_http(&issuer);
64 if issuer.scheme() != "https" && !allows_loopback_http {
65 return Err(OidcAuthError::InsecureIssuer);
66 }
67 if issuer.query().is_some() || issuer.fragment().is_some() {
68 return Err(OidcAuthError::InvalidIssuerComponents);
69 }
70
71 let audience = audience.trim().to_owned();
72 if audience.is_empty() {
73 return Err(OidcAuthError::EmptyAudience);
74 }
75
76 Ok(Self {
77 issuer,
78 issuer_value,
79 audience,
80 allows_loopback_http,
81 })
82 }
83
84 fn discovery_url(&self) -> Result<Url, OidcAuthError> {
85 Url::parse(&format!(
86 "{}/.well-known/openid-configuration",
87 self.issuer.as_str().trim_end_matches('/')
88 ))
89 .map_err(OidcAuthError::InvalidIssuer)
90 }
91}
92
93fn is_loopback_http(url: &Url) -> bool {
94 url.scheme() == "http"
95 && match url.host() {
96 Some(Host::Ipv4(address)) => address.is_loopback(),
97 Some(Host::Ipv6(address)) => address.is_loopback(),
98 Some(Host::Domain(_)) | None => false,
99 }
100}
101
102struct CachedKey {
103 decoding_key: DecodingKey,
104 algorithm: Option<Algorithm>,
105}
106
107struct CachedJwks {
108 keys: HashMap<String, Arc<CachedKey>>,
109 expires_at: Instant,
110}
111
112struct RefreshState {
113 last_completed: Instant,
114 retry_after: Option<Instant>,
115 coalesce_until: Option<Instant>,
116}
117
118struct RefreshAttempt {
119 refresh_state: Arc<std::sync::Mutex<RefreshState>>,
120 retry_after: Instant,
121 clear_on_drop: bool,
122}
123
124impl RefreshAttempt {
125 fn begin(refresh_state: Arc<std::sync::Mutex<RefreshState>>) -> Result<Self, OidcAuthError> {
126 let retry_after = Instant::now() + JWKS_REFRESH_COOLDOWN;
127 refresh_state
128 .lock()
129 .map_err(|_| OidcAuthError::RefreshStateUnavailable)?
130 .retry_after = Some(retry_after);
131 Ok(Self {
132 refresh_state,
133 retry_after,
134 clear_on_drop: true,
135 })
136 }
137
138 fn retain_backoff(mut self) {
139 self.clear_on_drop = false;
140 }
141
142 fn complete(mut self) {
143 self.clear_on_drop = false;
144 }
145}
146
147impl Drop for RefreshAttempt {
148 fn drop(&mut self) {
149 if !self.clear_on_drop {
150 return;
151 }
152 if let Ok(mut refresh_state) = self.refresh_state.lock()
153 && refresh_state.retry_after == Some(self.retry_after)
154 {
155 refresh_state.retry_after = None;
156 }
157 }
158}
159
160#[derive(Clone)]
161pub struct OidcAuthenticator {
162 audience: String,
163 jwks_uri: Url,
164 keys: Arc<RwLock<CachedJwks>>,
165 refresh_state: Arc<std::sync::Mutex<RefreshState>>,
166 refresh_gate: Arc<tokio::sync::Semaphore>,
167 validations: Arc<Vec<(Algorithm, Validation)>>,
168 client: reqwest::Client,
169}
170
171impl OidcAuthenticator {
172 pub async fn discover(config: OidcConfig) -> Result<Self, OidcAuthError> {
179 let client = reqwest::Client::builder()
180 .redirect(reqwest::redirect::Policy::none())
181 .timeout(OIDC_HTTP_TIMEOUT)
182 .build()
183 .map_err(OidcAuthError::HttpClient)?;
184 let (metadata, _) =
185 fetch_json::<ProviderMetadata>(&client, config.discovery_url()?, ProviderRequest::Metadata).await?;
186 let discovered_issuer = Url::parse(&metadata.issuer).map_err(OidcAuthError::InvalidDiscoveredIssuer)?;
187 if discovered_issuer != config.issuer {
188 return Err(OidcAuthError::IssuerMismatch {
189 expected: config.issuer_value,
190 discovered: metadata.issuer,
191 });
192 }
193
194 let jwks_uri = Url::parse(&metadata.jwks_uri).map_err(OidcAuthError::InvalidJwksUri)?;
195 if jwks_uri.scheme() != "https" && !(config.allows_loopback_http && is_loopback_http(&jwks_uri)) {
196 return Err(OidcAuthError::InsecureJwksUri);
197 }
198
199 let keys = fetch_jwks(&client, jwks_uri.clone()).await?;
200 let refresh_completed = Instant::now();
201 let validations = build_validations(&metadata.issuer, &config.audience);
202 Ok(Self {
203 audience: config.audience,
204 jwks_uri,
205 keys: Arc::new(RwLock::new(keys)),
206 refresh_state: Arc::new(std::sync::Mutex::new(RefreshState {
207 last_completed: refresh_completed,
208 retry_after: None,
209 coalesce_until: None,
210 })),
211 refresh_gate: Arc::new(tokio::sync::Semaphore::new(1)),
212 validations: Arc::new(validations),
213 client,
214 })
215 }
216
217 async fn authenticate(&self, token: &str) -> Result<AuthenticatedPrincipal, OidcAuthError> {
218 let token_header = decode_header(token).map_err(OidcAuthError::InvalidToken)?;
219 if matches!(token_header.alg, Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512) {
220 return Err(OidcAuthError::UnsupportedTokenAlgorithm);
221 }
222 let kid = token_header.kid.ok_or(OidcAuthError::MissingKeyId)?;
223 let key = self.verification_key(&kid).await?;
224 if key.algorithm.is_some_and(|algorithm| algorithm != token_header.alg) {
225 return Err(OidcAuthError::AlgorithmMismatch);
226 }
227 let validation = self
228 .validations
229 .iter()
230 .find_map(|(algorithm, validation)| (*algorithm == token_header.alg).then_some(validation))
231 .ok_or(OidcAuthError::UnsupportedTokenAlgorithm)?;
232
233 let claims = decode::<IdentityClaims>(token, &key.decoding_key, validation)
234 .map_err(OidcAuthError::InvalidToken)?
235 .claims;
236 if claims.sub.is_empty() {
237 return Err(OidcAuthError::EmptySubject);
238 }
239 if !claims.audience_allows(&self.audience) {
240 return Err(OidcAuthError::InvalidAuthorizedParty);
241 }
242 Ok(AuthenticatedPrincipal {
243 issuer: claims.iss,
244 subject: claims.sub,
245 expires_at: claims.exp,
246 })
247 }
248
249 async fn verification_key(&self, kid: &str) -> Result<Arc<CachedKey>, OidcAuthError> {
250 {
251 let keys = self.keys.read().await;
252 if Instant::now() < keys.expires_at {
253 if let Some(key) = keys.keys.get(kid) {
254 return Ok(Arc::clone(key));
255 }
256 }
257 }
258
259 let _refresh_permit = tokio::time::timeout(JWKS_REFRESH_WAIT_TIMEOUT, self.refresh_gate.acquire())
260 .await
261 .map_err(|_| OidcAuthError::JwksRefreshBackoff)?
262 .map_err(|_| OidcAuthError::RefreshStateUnavailable)?;
263 {
264 let keys = self.keys.read().await;
265 let refresh_state = self
266 .refresh_state
267 .lock()
268 .map_err(|_| OidcAuthError::RefreshStateUnavailable)?;
269 let now = Instant::now();
270 if now < keys.expires_at {
271 if let Some(key) = keys.keys.get(kid) {
272 return Ok(Arc::clone(key));
273 }
274 if now.duration_since(refresh_state.last_completed) < JWKS_REFRESH_COOLDOWN {
275 return Err(OidcAuthError::UnknownKeyId);
276 }
277 }
278 if refresh_state.coalesce_until.is_some_and(|deadline| now < deadline) {
279 return keys.keys.get(kid).cloned().ok_or(OidcAuthError::UnknownKeyId);
280 }
281 if refresh_state.retry_after.is_some_and(|deadline| now < deadline) {
282 return Err(OidcAuthError::JwksRefreshBackoff);
283 }
284 }
285
286 let refresh_attempt = RefreshAttempt::begin(Arc::clone(&self.refresh_state))?;
287 let refreshed = match fetch_jwks(&self.client, self.jwks_uri.clone()).await {
288 Ok(refreshed) => refreshed,
289 Err(error) => {
290 refresh_attempt.retain_backoff();
291 return Err(error);
292 }
293 };
294 let key = refreshed.keys.get(kid).cloned();
295 *self.keys.write().await = refreshed;
296 let refresh_completed = Instant::now();
297 let mut refresh_state = self
298 .refresh_state
299 .lock()
300 .map_err(|_| OidcAuthError::RefreshStateUnavailable)?;
301 refresh_state.last_completed = refresh_completed;
302 refresh_state.retry_after = None;
303 refresh_state.coalesce_until = Some(refresh_completed + JWKS_REFRESH_COALESCE_WINDOW);
304 drop(refresh_state);
305 refresh_attempt.complete();
306 key.ok_or(OidcAuthError::UnknownKeyId)
307 }
308}
309
310#[derive(Clone, Debug, Eq, PartialEq)]
311pub struct AuthenticatedPrincipal {
312 issuer: String,
313 subject: String,
314 expires_at: u64,
315}
316
317impl AuthenticatedPrincipal {
318 #[must_use]
319 pub fn issuer(&self) -> &str {
320 &self.issuer
321 }
322
323 #[must_use]
324 pub fn subject(&self) -> &str {
325 &self.subject
326 }
327
328 #[must_use]
329 pub fn expires_at(&self) -> u64 {
330 self.expires_at
331 }
332
333 pub(crate) fn is_expired(&self) -> bool {
334 self.is_expired_at(jsonwebtoken::get_current_timestamp())
335 }
336
337 fn is_expired_at(&self, timestamp: u64) -> bool {
338 timestamp > self.expires_at.saturating_add(JWT_CLOCK_SKEW_SECONDS)
339 }
340
341 #[cfg(test)]
342 pub(crate) fn expired_for_test() -> Self {
343 Self {
344 issuer: "https://issuer.example".to_owned(),
345 subject: "subject".to_owned(),
346 expires_at: 0,
347 }
348 }
349}
350
351pub async fn require_oidc(
352 State(authenticator): State<OidcAuthenticator>,
353 mut request: Request<Body>,
354 next: Next,
355) -> Response {
356 let error_format = AuthErrorFormat::for_path(request.uri().path());
357 let Some(token) = bearer_token(request.headers()) else {
358 return authentication_error(error_format, "missing_bearer_token", "missing bearer token");
359 };
360 let duplicate_identity_api_key = request
361 .headers()
362 .get_all("x-api-key")
363 .iter()
364 .any(|value| value.to_str().is_ok_and(|value| value.trim() == token));
365
366 match authenticator.authenticate(token).await {
367 Ok(principal) => {
368 request.headers_mut().remove(header::AUTHORIZATION);
369 if duplicate_identity_api_key || !error_format.allows_upstream_api_key() {
370 request.headers_mut().remove("x-api-key");
371 }
372 request.extensions_mut().insert(principal);
373 next.run(request).await
374 }
375 Err(error) => {
376 if error.is_dependency_failure() {
377 warn!(error = %error, "OIDC token verification dependency failed");
378 authentication_service_unavailable(error_format)
379 } else {
380 debug!(error = %error, "OIDC bearer token rejected");
381 authentication_error(error_format, "invalid_token", "invalid bearer token")
382 }
383 }
384 }
385}
386
387fn bearer_token(headers: &axum::http::HeaderMap) -> Option<&str> {
388 headers
389 .get(header::AUTHORIZATION)?
390 .to_str()
391 .ok()?
392 .split_once(' ')
393 .and_then(|(scheme, token)| {
394 let token = token.trim();
395 (scheme.eq_ignore_ascii_case("bearer") && !token.is_empty()).then_some(token)
396 })
397}
398
399#[derive(Clone, Copy)]
400enum AuthErrorFormat {
401 OpenAi,
402 Anthropic,
403}
404
405impl AuthErrorFormat {
406 fn for_path(path: &str) -> Self {
407 if matches!(path, ANTHROPIC_MESSAGES_PATH | ANTHROPIC_COUNT_TOKENS_PATH) {
408 Self::Anthropic
409 } else {
410 Self::OpenAi
411 }
412 }
413
414 fn allows_upstream_api_key(self) -> bool {
415 matches!(self, Self::Anthropic)
416 }
417}
418
419fn authentication_error(format: AuthErrorFormat, code: &'static str, message: &'static str) -> Response {
420 protocol_error(
421 format,
422 StatusCode::UNAUTHORIZED,
423 "authentication_error",
424 "authentication_error",
425 code,
426 message,
427 true,
428 )
429}
430
431fn authentication_service_unavailable(format: AuthErrorFormat) -> Response {
432 protocol_error(
433 format,
434 StatusCode::SERVICE_UNAVAILABLE,
435 "server_error",
436 "api_error",
437 "authentication_service_unavailable",
438 "authentication service temporarily unavailable",
439 false,
440 )
441}
442
443fn protocol_error(
444 format: AuthErrorFormat,
445 status: StatusCode,
446 openai_error_type: &'static str,
447 anthropic_error_type: &'static str,
448 code: &'static str,
449 message: &'static str,
450 challenge: bool,
451) -> Response {
452 let (body, request_id) = match format {
453 AuthErrorFormat::OpenAi => (
454 json!({
455 "error": {
456 "message": message,
457 "type": openai_error_type,
458 "param": null,
459 "code": code
460 }
461 }),
462 None,
463 ),
464 AuthErrorFormat::Anthropic => {
465 let request_id = uuid7_str("req_");
466 (
467 json!({
468 "type": "error",
469 "error": {
470 "type": anthropic_error_type,
471 "message": message
472 },
473 "request_id": &request_id
474 }),
475 Some(request_id),
476 )
477 }
478 };
479 let mut builder = Response::builder()
480 .status(status)
481 .header(header::CONTENT_TYPE, "application/json");
482 if let Some(request_id) = request_id.as_deref() {
483 builder = builder.header("request-id", request_id);
484 }
485 if challenge {
486 builder = builder.header(header::WWW_AUTHENTICATE, "Bearer");
487 }
488 builder
489 .body(Body::from(body.to_string()))
490 .expect("valid authentication protocol error response")
491}
492
493#[derive(Clone, Copy)]
494enum ProviderRequest {
495 Metadata,
496 Jwks,
497}
498
499impl ProviderRequest {
500 fn error(self, error: reqwest::Error) -> OidcAuthError {
501 match self {
502 Self::Metadata => OidcAuthError::ProviderMetadataRequest(error),
503 Self::Jwks => OidcAuthError::JwksRequest(error),
504 }
505 }
506}
507
508async fn fetch_jwks(client: &reqwest::Client, uri: Url) -> Result<CachedJwks, OidcAuthError> {
509 let (keys, headers) = fetch_json::<JwkSet>(client, uri, ProviderRequest::Jwks).await?;
510 compile_jwks(keys, jwks_ttl(&headers))
511}
512
513async fn fetch_json<T>(
514 client: &reqwest::Client,
515 uri: Url,
516 request_kind: ProviderRequest,
517) -> Result<(T, HeaderMap), OidcAuthError>
518where
519 T: serde::de::DeserializeOwned,
520{
521 let mut response = client
522 .get(uri)
523 .send()
524 .await
525 .map_err(|error| request_kind.error(error))?
526 .error_for_status()
527 .map_err(|error| request_kind.error(error))?;
528 if response
529 .content_length()
530 .is_some_and(|length| length > MAX_PROVIDER_RESPONSE_BYTES as u64)
531 {
532 return Err(OidcAuthError::ProviderResponseTooLarge);
533 }
534
535 let headers = response.headers().clone();
536 let mut body = Vec::with_capacity(
537 response
538 .content_length()
539 .and_then(|length| usize::try_from(length).ok())
540 .unwrap_or_default()
541 .min(MAX_PROVIDER_RESPONSE_BYTES),
542 );
543 while let Some(chunk) = response.chunk().await.map_err(|error| request_kind.error(error))? {
544 if body.len().saturating_add(chunk.len()) > MAX_PROVIDER_RESPONSE_BYTES {
545 return Err(OidcAuthError::ProviderResponseTooLarge);
546 }
547 body.extend_from_slice(&chunk);
548 }
549
550 let value = serde_json::from_slice(&body).map_err(OidcAuthError::InvalidProviderJson)?;
551 Ok((value, headers))
552}
553
554fn compile_jwks(keys: JwkSet, ttl: Duration) -> Result<CachedJwks, OidcAuthError> {
555 if keys.keys.is_empty() {
556 return Err(OidcAuthError::EmptyJwks);
557 }
558 if keys.keys.len() > MAX_JWKS_KEYS {
559 return Err(OidcAuthError::TooManyJwksKeys);
560 }
561
562 let mut compiled = HashMap::with_capacity(keys.keys.len());
563 for key in keys.keys {
564 let Some(kid) = key.common.key_id.clone() else {
565 continue;
566 };
567 let algorithm = match verification_algorithm(&key) {
568 VerificationAlgorithm::Skip => continue,
569 VerificationAlgorithm::AnyAsymmetric => None,
570 VerificationAlgorithm::Exact(algorithm) => Some(algorithm),
571 };
572 let Ok(decoding_key) = DecodingKey::from_jwk(&key) else {
573 continue;
574 };
575 if compiled
576 .insert(
577 kid.clone(),
578 Arc::new(CachedKey {
579 decoding_key,
580 algorithm,
581 }),
582 )
583 .is_some()
584 {
585 return Err(OidcAuthError::DuplicateKeyId(kid));
586 }
587 }
588 if compiled.is_empty() {
589 return Err(OidcAuthError::EmptyJwks);
590 }
591
592 Ok(CachedJwks {
593 keys: compiled,
594 expires_at: Instant::now() + ttl,
595 })
596}
597
598enum VerificationAlgorithm {
599 Skip,
600 AnyAsymmetric,
601 Exact(Algorithm),
602}
603
604fn verification_algorithm(key: &Jwk) -> VerificationAlgorithm {
605 if key
606 .common
607 .public_key_use
608 .as_ref()
609 .is_some_and(|key_use| key_use != &PublicKeyUse::Signature)
610 || key
611 .common
612 .key_operations
613 .as_ref()
614 .is_some_and(|operations| !operations.contains(&KeyOperations::Verify))
615 {
616 return VerificationAlgorithm::Skip;
617 }
618
619 match key.common.key_algorithm {
620 Some(key_algorithm) => {
621 let algorithm = match key_algorithm {
622 KeyAlgorithm::ES256 => Algorithm::ES256,
623 KeyAlgorithm::ES384 => Algorithm::ES384,
624 KeyAlgorithm::RS256 => Algorithm::RS256,
625 KeyAlgorithm::RS384 => Algorithm::RS384,
626 KeyAlgorithm::RS512 => Algorithm::RS512,
627 KeyAlgorithm::PS256 => Algorithm::PS256,
628 KeyAlgorithm::PS384 => Algorithm::PS384,
629 KeyAlgorithm::PS512 => Algorithm::PS512,
630 KeyAlgorithm::EdDSA => Algorithm::EdDSA,
631 KeyAlgorithm::HS256
632 | KeyAlgorithm::HS384
633 | KeyAlgorithm::HS512
634 | KeyAlgorithm::RSA1_5
635 | KeyAlgorithm::RSA_OAEP
636 | KeyAlgorithm::RSA_OAEP_256
637 | KeyAlgorithm::UNKNOWN_ALGORITHM => return VerificationAlgorithm::Skip,
638 };
639 if SUPPORTED_VERIFICATION_ALGORITHMS.contains(&algorithm) {
640 VerificationAlgorithm::Exact(algorithm)
641 } else {
642 VerificationAlgorithm::Skip
643 }
644 }
645 None => VerificationAlgorithm::AnyAsymmetric,
646 }
647}
648
649fn build_validations(issuer: &str, audience: &str) -> Vec<(Algorithm, Validation)> {
650 SUPPORTED_VERIFICATION_ALGORITHMS
651 .into_iter()
652 .map(|algorithm| {
653 let mut validation = Validation::new(algorithm);
654 validation.leeway = JWT_CLOCK_SKEW_SECONDS;
655 validation.set_audience(&[audience]);
656 validation.set_issuer(&[issuer]);
657 validation.set_required_spec_claims(&["exp", "iss", "aud", "sub"]);
658 validation.validate_nbf = true;
659 (algorithm, validation)
660 })
661 .collect()
662}
663
664fn jwks_ttl(headers: &HeaderMap) -> Duration {
665 let mut max_age: Option<u64> = None;
666 for value in headers.get_all(header::CACHE_CONTROL) {
667 let Ok(value) = value.to_str() else {
668 continue;
669 };
670 for directive in value.split(',').map(str::trim) {
671 if directive.eq_ignore_ascii_case("no-cache") || directive.eq_ignore_ascii_case("no-store") {
672 return Duration::ZERO;
673 }
674 if let Some((name, seconds)) = directive.split_once('=') {
675 if name.trim().eq_ignore_ascii_case("max-age") {
676 let parsed = seconds.trim().trim_matches('"').parse::<u64>().ok();
677 max_age = match (max_age, parsed) {
678 (Some(existing), Some(parsed)) => Some(existing.min(parsed)),
679 (None, parsed) => parsed,
680 (existing, None) => existing,
681 };
682 }
683 }
684 }
685 }
686 let age = headers
687 .get_all(header::AGE)
688 .iter()
689 .filter_map(|value| value.to_str().ok()?.trim().parse::<u64>().ok())
690 .max()
691 .map_or(Duration::ZERO, Duration::from_secs);
692 max_age.map_or_else(
693 || DEFAULT_JWKS_TTL.saturating_sub(age),
694 |seconds| Duration::from_secs(seconds).saturating_sub(age).min(MAX_JWKS_TTL),
695 )
696}
697
698#[derive(Deserialize)]
699struct ProviderMetadata {
700 issuer: String,
701 jwks_uri: String,
702}
703
704#[derive(Deserialize)]
705struct IdentityClaims {
706 iss: String,
707 sub: String,
708 aud: AudienceClaim,
709 exp: u64,
710 #[serde(default)]
711 azp: Option<String>,
712}
713
714impl IdentityClaims {
715 fn audience_allows(&self, expected: &str) -> bool {
716 let audiences_match = match &self.aud {
717 AudienceClaim::One(audience) => audience == expected,
718 AudienceClaim::Many(audiences) => {
719 !audiences.is_empty() && audiences.iter().all(|audience| audience == expected)
720 }
721 };
722 audiences_match
723 && self
724 .azp
725 .as_deref()
726 .is_none_or(|authorized_party| authorized_party == expected)
727 }
728}
729
730#[derive(Deserialize)]
731#[serde(untagged)]
732enum AudienceClaim {
733 One(String),
734 Many(Vec<String>),
735}
736
737#[derive(Debug, thiserror::Error)]
738#[non_exhaustive]
739pub enum OidcAuthError {
740 #[error("OIDC audience must not be empty")]
741 EmptyAudience,
742 #[error("OIDC issuer must not include a query or fragment")]
743 InvalidIssuerComponents,
744 #[error("OIDC issuer must use HTTPS; HTTP is only permitted for loopback IP addresses (127.0.0.1 or ::1)")]
745 InsecureIssuer,
746 #[error("OIDC JWKS URI must use HTTPS; HTTP is only permitted for loopback IP addresses (127.0.0.1 or ::1)")]
747 InsecureJwksUri,
748 #[error("invalid OIDC issuer URL")]
749 InvalidIssuer(#[source] url::ParseError),
750 #[error("invalid OIDC JWKS URI")]
751 InvalidJwksUri(#[source] url::ParseError),
752 #[error("OIDC discovery returned an invalid issuer URL")]
753 InvalidDiscoveredIssuer(#[source] url::ParseError),
754 #[error("failed to build OIDC HTTP client")]
755 HttpClient(#[source] reqwest::Error),
756 #[error("OIDC provider metadata request failed")]
757 ProviderMetadataRequest(#[source] reqwest::Error),
758 #[error("OIDC JWKS request failed")]
759 JwksRequest(#[source] reqwest::Error),
760 #[error("OIDC provider response exceeded the size limit")]
761 ProviderResponseTooLarge,
762 #[error("OIDC provider returned invalid JSON")]
763 InvalidProviderJson(#[source] serde_json::Error),
764 #[error("OIDC discovery returned issuer {discovered}, expected {expected}")]
765 IssuerMismatch { expected: String, discovered: String },
766 #[error("OIDC provider returned an empty JWKS")]
767 EmptyJwks,
768 #[error("OIDC provider returned too many JWKs")]
769 TooManyJwksKeys,
770 #[error("OIDC provider returned duplicate JWK key ID {0}")]
771 DuplicateKeyId(String),
772 #[error("OIDC JWKS refresh is temporarily backed off after a provider failure")]
773 JwksRefreshBackoff,
774 #[error("OIDC JWKS refresh coordination is unavailable")]
775 RefreshStateUnavailable,
776 #[error("bearer token is missing a key ID")]
777 MissingKeyId,
778 #[error("bearer token references an unknown key ID")]
779 UnknownKeyId,
780 #[error("bearer token subject must not be empty")]
781 EmptySubject,
782 #[error("bearer token authorized party does not match the configured audience")]
783 InvalidAuthorizedParty,
784 #[error("bearer token uses an unsupported algorithm")]
785 UnsupportedTokenAlgorithm,
786 #[error("bearer token and JWK algorithms do not match")]
787 AlgorithmMismatch,
788 #[error("bearer token validation failed")]
789 InvalidToken(#[source] jsonwebtoken::errors::Error),
790}
791
792impl OidcAuthError {
793 fn is_dependency_failure(&self) -> bool {
794 matches!(
795 self,
796 Self::JwksRequest(_)
797 | Self::ProviderResponseTooLarge
798 | Self::InvalidProviderJson(_)
799 | Self::EmptyJwks
800 | Self::TooManyJwksKeys
801 | Self::DuplicateKeyId(_)
802 | Self::JwksRefreshBackoff
803 | Self::RefreshStateUnavailable
804 )
805 }
806}
807
808#[cfg(test)]
809mod tests {
810 use super::{
811 AuthenticatedPrincipal, JWKS_REFRESH_COOLDOWN, MAX_JWKS_KEYS, MAX_JWKS_TTL, OidcAuthError, OidcAuthenticator,
812 OidcConfig, VerificationAlgorithm, compile_jwks, jwks_ttl, verification_algorithm,
813 };
814 use axum::http::{HeaderMap, HeaderValue, header};
815 use axum::routing::get;
816 use axum::{Json, Router};
817 use jsonwebtoken::jwk::{Jwk, JwkSet, KeyAlgorithm, KeyOperations, PublicKeyUse};
818 use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
819 use rand::rngs::OsRng;
820 use rsa::RsaPrivateKey;
821 use rsa::pkcs1::EncodeRsaPrivateKey;
822 use serde_json::json;
823 use std::sync::Arc;
824 use std::sync::atomic::{AtomicUsize, Ordering};
825 use std::time::{Duration, Instant};
826 use tokio::net::TcpListener;
827 use tokio::sync::Semaphore;
828
829 fn test_jwk() -> Jwk {
830 let private_key = RsaPrivateKey::new(&mut OsRng, 2048).expect("generate test RSA key");
831 let private_key = private_key.to_pkcs1_der().expect("encode test RSA key");
832 let mut jwk = Jwk::from_encoding_key(&EncodingKey::from_rsa_der(private_key.as_bytes()), Algorithm::RS256)
833 .expect("test JWK");
834 jwk.common.key_id = Some("test-key".to_owned());
835 jwk.common.key_algorithm = Some(KeyAlgorithm::RS256);
836 jwk.common.public_key_use = Some(PublicKeyUse::Signature);
837 jwk
838 }
839
840 fn test_key_with_id(kid: &str) -> (Vec<u8>, Jwk) {
841 let private_key = RsaPrivateKey::new(&mut OsRng, 2048).expect("generate test RSA key");
842 let private_key = private_key.to_pkcs1_der().expect("encode test RSA key");
843 let private_key = private_key.as_bytes().to_vec();
844 let mut jwk =
845 Jwk::from_encoding_key(&EncodingKey::from_rsa_der(&private_key), Algorithm::RS256).expect("test JWK");
846 jwk.common.key_id = Some(kid.to_owned());
847 jwk.common.key_algorithm = Some(KeyAlgorithm::RS256);
848 jwk.common.public_key_use = Some(PublicKeyUse::Signature);
849 (private_key, jwk)
850 }
851
852 async fn cancellation_test_authenticator() -> (
853 OidcAuthenticator,
854 String,
855 Arc<AtomicUsize>,
856 Arc<Semaphore>,
857 Arc<Semaphore>,
858 tokio::task::JoinHandle<()>,
859 ) {
860 let listener = TcpListener::bind("127.0.0.1:0")
861 .await
862 .expect("bind cancellation test provider");
863 let issuer = format!("http://{}", listener.local_addr().expect("provider address"));
864 let discovery_issuer = issuer.clone();
865 let discovery_jwks_uri = format!("{issuer}/jwks");
866 let (_old_private_key, old_jwk) = test_key_with_id("old-key");
867 let (new_private_key, new_jwk) = test_key_with_id("new-key");
868 let requests = Arc::new(AtomicUsize::new(0));
869 let observed_requests = Arc::clone(&requests);
870 let refresh_started = Arc::new(Semaphore::new(0));
871 let observed_refresh_started = Arc::clone(&refresh_started);
872 let release_refresh = Arc::new(Semaphore::new(0));
873 let observed_release_refresh = Arc::clone(&release_refresh);
874
875 let provider = Router::new()
876 .route(
877 "/.well-known/openid-configuration",
878 get(move || {
879 let issuer = discovery_issuer.clone();
880 let jwks_uri = discovery_jwks_uri.clone();
881 async move { Json(json!({"issuer": issuer, "jwks_uri": jwks_uri})) }
882 }),
883 )
884 .route(
885 "/jwks",
886 get(move || {
887 let old_jwk = old_jwk.clone();
888 let new_jwk = new_jwk.clone();
889 let requests = Arc::clone(&observed_requests);
890 let refresh_started = Arc::clone(&observed_refresh_started);
891 let release_refresh = Arc::clone(&observed_release_refresh);
892 async move {
893 let request = requests.fetch_add(1, Ordering::Relaxed);
894 let jwk = if request == 0 {
895 old_jwk
896 } else {
897 refresh_started.add_permits(1);
898 release_refresh
899 .acquire()
900 .await
901 .expect("release semaphore remains open")
902 .forget();
903 new_jwk
904 };
905 ([(header::CACHE_CONTROL, "max-age=0")], Json(json!({"keys": [jwk]})))
906 }
907 }),
908 );
909 let provider_handle = tokio::spawn(async move {
910 axum::serve(listener, provider)
911 .await
912 .expect("serve cancellation test provider");
913 });
914 let authenticator = OidcAuthenticator::discover(OidcConfig::new(&issuer, "agentic-api").expect("OIDC config"))
915 .await
916 .expect("discover cancellation test provider");
917 let mut token_header = Header::new(Algorithm::RS256);
918 token_header.kid = Some("new-key".to_owned());
919 let token = encode(
920 &token_header,
921 &json!({
922 "iss": issuer,
923 "sub": "subject",
924 "aud": "agentic-api",
925 "exp": jsonwebtoken::get_current_timestamp() + 300
926 }),
927 &EncodingKey::from_rsa_der(&new_private_key),
928 )
929 .expect("encode cancellation test token");
930
931 (
932 authenticator,
933 token,
934 requests,
935 refresh_started,
936 release_refresh,
937 provider_handle,
938 )
939 }
940
941 #[test]
942 fn jwks_cache_lifetime_uses_provider_max_age_with_a_cap() {
943 let mut headers = HeaderMap::new();
944 headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("public, max-age=60"));
945 assert_eq!(jwks_ttl(&headers), Duration::from_secs(60));
946
947 headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("max-age=86400"));
948 assert_eq!(jwks_ttl(&headers), MAX_JWKS_TTL);
949
950 headers.insert(
951 header::CACHE_CONTROL,
952 HeaderValue::from_static("private, Max-Age=\"30\""),
953 );
954 assert_eq!(jwks_ttl(&headers), Duration::from_secs(30));
955
956 headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
957 assert_eq!(jwks_ttl(&headers), Duration::ZERO);
958
959 headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("max-age=60"));
960 headers.insert(header::AGE, HeaderValue::from_static("55"));
961 assert_eq!(jwks_ttl(&headers), Duration::from_secs(5));
962
963 headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("max-age=86400"));
964 headers.insert(header::AGE, HeaderValue::from_static("4000"));
965 assert_eq!(jwks_ttl(&headers), MAX_JWKS_TTL);
966
967 headers.remove(header::AGE);
968 headers.append(header::CACHE_CONTROL, HeaderValue::from_static("no-cache"));
969 assert_eq!(jwks_ttl(&headers), Duration::ZERO);
970 }
971
972 #[test]
973 fn authenticated_principal_expiration_includes_clock_skew() {
974 let principal = AuthenticatedPrincipal {
975 issuer: "https://issuer.example".to_owned(),
976 subject: "subject".to_owned(),
977 expires_at: 100,
978 };
979
980 assert!(!principal.is_expired_at(160));
981 assert!(principal.is_expired_at(161));
982 }
983
984 #[tokio::test]
985 async fn cancelled_jwks_refresh_does_not_install_backoff() {
986 let (authenticator, token, requests, refresh_started, release_refresh, _provider) =
987 cancellation_test_authenticator().await;
988 let cancelled_authenticator = authenticator.clone();
989 let cancelled_token = token.clone();
990 let cancelled_refresh =
991 tokio::spawn(async move { cancelled_authenticator.authenticate(&cancelled_token).await });
992
993 tokio::time::timeout(Duration::from_secs(2), refresh_started.acquire())
994 .await
995 .expect("refresh must start")
996 .expect("refresh semaphore remains open")
997 .forget();
998 cancelled_refresh.abort();
999 assert!(
1000 cancelled_refresh
1001 .await
1002 .expect_err("refresh task must be cancelled")
1003 .is_cancelled()
1004 );
1005 release_refresh.add_permits(2);
1006
1007 let principal = tokio::time::timeout(Duration::from_secs(2), authenticator.authenticate(&token))
1008 .await
1009 .expect("replacement refresh must finish")
1010 .expect("replacement refresh must authenticate");
1011 assert_eq!(principal.subject(), "subject");
1012 assert_eq!(requests.load(Ordering::Relaxed), 3);
1013 }
1014
1015 #[tokio::test]
1016 async fn stalled_jwks_refresh_bounds_waiters() {
1017 let (authenticator, token, _requests, refresh_started, release_refresh, _provider) =
1018 cancellation_test_authenticator().await;
1019 let refreshing_authenticator = authenticator.clone();
1020 let refreshing_token = token.clone();
1021 let refreshing = tokio::spawn(async move { refreshing_authenticator.authenticate(&refreshing_token).await });
1022 tokio::time::timeout(Duration::from_secs(2), refresh_started.acquire())
1023 .await
1024 .expect("refresh must start")
1025 .expect("refresh semaphore remains open")
1026 .forget();
1027
1028 let waiting = tokio::time::timeout(Duration::from_secs(2), authenticator.authenticate(&token))
1029 .await
1030 .expect("refresh waiter must be bounded");
1031 assert!(matches!(waiting, Err(OidcAuthError::JwksRefreshBackoff)));
1032
1033 refreshing.abort();
1034 assert!(
1035 refreshing
1036 .await
1037 .expect_err("refresh task must be cancelled")
1038 .is_cancelled()
1039 );
1040 release_refresh.add_permits(1);
1041 }
1042
1043 #[tokio::test]
1044 async fn fresh_cache_unknown_key_refreshes_once_after_cooldown() {
1045 let (authenticator, token, requests, _refresh_started, release_refresh, _provider) =
1046 cancellation_test_authenticator().await;
1047 authenticator.keys.write().await.expires_at = Instant::now() + Duration::from_secs(60);
1048 authenticator
1049 .refresh_state
1050 .lock()
1051 .expect("refresh state")
1052 .last_completed = Instant::now()
1053 .checked_sub(JWKS_REFRESH_COOLDOWN + Duration::from_secs(1))
1054 .expect("test cooldown instant");
1055 release_refresh.add_permits(1);
1056
1057 let first_authenticator = authenticator.clone();
1058 let first_token = token.clone();
1059 let first = tokio::spawn(async move { first_authenticator.authenticate(&first_token).await });
1060 let second_authenticator = authenticator.clone();
1061 let second = tokio::spawn(async move { second_authenticator.authenticate(&token).await });
1062 let (first, second) = tokio::join!(first, second);
1063
1064 assert_eq!(
1065 first.expect("first task").expect("first authentication").subject(),
1066 "subject"
1067 );
1068 assert_eq!(
1069 second.expect("second task").expect("second authentication").subject(),
1070 "subject"
1071 );
1072 assert_eq!(requests.load(Ordering::Relaxed), 2);
1073 }
1074
1075 #[test]
1076 fn jwks_limits_and_signature_metadata_are_enforced() {
1077 let mut encryption_key = test_jwk();
1078 encryption_key.common.public_key_use = Some(PublicKeyUse::Encryption);
1079 assert!(matches!(
1080 compile_jwks(
1081 JwkSet {
1082 keys: vec![encryption_key]
1083 },
1084 Duration::from_secs(60)
1085 ),
1086 Err(OidcAuthError::EmptyJwks)
1087 ));
1088
1089 let mut non_verifying_key = test_jwk();
1090 non_verifying_key.common.key_operations = Some(vec![KeyOperations::Encrypt]);
1091 assert!(matches!(
1092 compile_jwks(
1093 JwkSet {
1094 keys: vec![non_verifying_key]
1095 },
1096 Duration::from_secs(60)
1097 ),
1098 Err(OidcAuthError::EmptyJwks)
1099 ));
1100
1101 let mut mismatched_algorithm = test_jwk();
1102 mismatched_algorithm.common.key_algorithm = Some(KeyAlgorithm::RS512);
1103 assert!(matches!(
1104 verification_algorithm(&mismatched_algorithm),
1105 VerificationAlgorithm::Exact(Algorithm::RS512)
1106 ));
1107
1108 let too_many_keys = vec![test_jwk(); MAX_JWKS_KEYS + 1];
1109 assert!(matches!(
1110 compile_jwks(JwkSet { keys: too_many_keys }, Duration::from_secs(60)),
1111 Err(OidcAuthError::TooManyJwksKeys)
1112 ));
1113
1114 let duplicate_key = test_jwk();
1115 let duplicate_key_id = duplicate_key.common.key_id.clone().expect("test key ID");
1116 assert!(matches!(
1117 compile_jwks(
1118 JwkSet {
1119 keys: vec![duplicate_key.clone(), duplicate_key]
1120 },
1121 Duration::from_secs(60)
1122 ),
1123 Err(OidcAuthError::DuplicateKeyId(key_id)) if key_id == duplicate_key_id
1124 ));
1125 }
1126}