Skip to main content

agentic_server/
auth.rs

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    /// Create an OIDC bearer-token configuration.
55    ///
56    /// # Errors
57    ///
58    /// Returns an error when the issuer is not an absolute HTTPS URL (except
59    /// for loopback test/development issuers) or the audience is empty.
60    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    /// Discover an OIDC provider and cache its initial verification keys.
173    ///
174    /// # Errors
175    ///
176    /// Returns an error when provider discovery or the JSON Web Key Set
177    /// (JWKS) request fails, or when the discovered metadata is inconsistent.
178    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}