Skip to main content

qcs_api_client_common/configuration/
tokens.rs

1//! Models and utilities for managing `OAuth2` sessions.
2use std::{pin::Pin, sync::Arc};
3
4use futures::Future;
5use jsonwebtoken::{Algorithm, DecodingKey, Validation};
6use oauth2::TokenResponse;
7use serde::{Deserialize, Serialize};
8use time::OffsetDateTime;
9use tokio::sync::{Mutex, Notify, RwLock};
10use tokio_util::sync::CancellationToken;
11
12#[cfg(feature = "stubs")]
13use pyo3_stub_gen::derive::gen_stub_pyclass;
14
15use super::{
16    ClientConfiguration, ConfigSource, TokenError, oidc, secrets::Secrets, settings::AuthServer,
17};
18use crate::configuration::{
19    error::{DiscoveryError, WriteError},
20    pkce::{PkceLoginError, PkceLoginRequest, RedirectBinding, pkce_login},
21    secrets::{Credential, SecretAccessToken, SecretRefreshToken, TokenPayload},
22};
23#[cfg(feature = "tracing-config")]
24use crate::tracing_configuration::TracingConfiguration;
25#[cfg(feature = "tracing")]
26use urlpattern::UrlPatternMatchInput;
27
28pub use super::secret_string::ClientSecret;
29
30/// A single type containing an access token and an associated refresh token.
31#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
32#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
33#[cfg_attr(
34    feature = "python",
35    pyo3::pyclass(
36        eq,
37        get_all,
38        set_all,
39        module = "qcs_api_client_common._qcs_api_client_common.configuration",
40        from_py_object
41    )
42)]
43pub struct RefreshToken {
44    /// The token used to refresh the access token.
45    pub refresh_token: SecretRefreshToken,
46}
47
48impl RefreshToken {
49    /// Create a new [`RefreshToken`] with the given refresh token.
50    #[must_use]
51    pub const fn new(refresh_token: SecretRefreshToken) -> Self {
52        Self { refresh_token }
53    }
54
55    /// Request and return a new access token from the given authorization server using this refresh token.
56    /// Updates the refresh token in-place if the authorization server returns a new one.
57    ///
58    /// # Errors
59    ///
60    /// See [`TokenError`]
61    pub async fn request_access_token(
62        &mut self,
63        auth_server: &AuthServer,
64    ) -> Result<SecretAccessToken, TokenError> {
65        if self.refresh_token.is_empty() {
66            return Err(TokenError::NoRefreshToken);
67        }
68
69        let client = default_http_client()?;
70        let token_url = oidc::fetch_discovery(&client, &auth_server.issuer)
71            .await?
72            .token_endpoint;
73        let data = TokenRefreshRequest::new(&auth_server.client_id, self.refresh_token.secret());
74        let resp = client.post(token_url).form(&data).send().await?;
75
76        // `error_for_status()` discards the response body, which is where OAuth2 servers put the
77        // actual reason a refresh was rejected (e.g. `invalid_grant`). Log it before converting to
78        // an opaque error, since callers otherwise have no way to tell "the refresh token is
79        // expired/revoked" apart from "a network blip happened" - both currently look identical
80        // and silently fall back to an interactive login.
81        if let Err(error) = resp.error_for_status_ref() {
82            #[cfg(feature = "tracing")]
83            {
84                let status = resp.status();
85                let body = resp.text().await.unwrap_or_default();
86                tracing::warn!(
87                    %status,
88                    %body,
89                    "the auth server rejected the refresh token request"
90                );
91            }
92            return Err(error.into());
93        }
94
95        let RefreshTokenResponse {
96            access_token,
97            refresh_token,
98        } = resp.json().await?;
99
100        if let Some(refresh_token) = refresh_token {
101            self.refresh_token = refresh_token;
102        }
103        Ok(access_token)
104    }
105}
106
107#[derive(Deserialize, Debug, Serialize)]
108pub(super) struct ClientCredentialsResponse {
109    pub(super) access_token: SecretAccessToken,
110}
111
112/// A pair of Client ID and Client Secret, used to request an OAuth Client Credentials Grant
113#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
114#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
115#[cfg_attr(
116    feature = "python",
117    pyo3::pyclass(
118        eq,
119        get_all,
120        frozen,
121        module = "qcs_api_client_common._qcs_api_client_common.configuration",
122        from_py_object
123    )
124)]
125pub struct ClientCredentials {
126    /// The client ID
127    pub client_id: String,
128    /// The client secret.
129    pub client_secret: ClientSecret,
130}
131
132impl ClientCredentials {
133    #[must_use]
134    /// Construct a new [`ClientCredentials`]
135    pub fn new(client_id: impl Into<String>, client_secret: impl Into<ClientSecret>) -> Self {
136        Self {
137            client_id: client_id.into(),
138            client_secret: client_secret.into(),
139        }
140    }
141
142    /// Get the client ID.
143    #[must_use]
144    pub fn client_id(&self) -> &str {
145        &self.client_id
146    }
147
148    /// Get the client secret.
149    #[must_use]
150    pub const fn client_secret(&self) -> &ClientSecret {
151        &self.client_secret
152    }
153
154    /// Request and return an access token from the given auth server using this set of client credentials.
155    ///
156    /// # Errors
157    ///
158    /// See [`TokenError`]
159    pub async fn request_access_token(
160        &self,
161        auth_server: &AuthServer,
162    ) -> Result<SecretAccessToken, TokenError> {
163        let request = ClientCredentialsRequest::new(None);
164        let client = default_http_client()?;
165
166        let url = oidc::fetch_discovery(&client, &auth_server.issuer)
167            .await?
168            .token_endpoint;
169        let ready_to_send = client
170            .post(url)
171            .basic_auth(&self.client_id, Some(&self.client_secret.secret()))
172            .form(&request);
173        let response = ready_to_send.send().await?;
174
175        response.error_for_status_ref()?;
176
177        let ClientCredentialsResponse { access_token } = response.json().await?;
178        Ok(access_token)
179    }
180}
181
182#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
183#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
184#[cfg_attr(
185    feature = "python",
186    pyo3::pyclass(
187        eq,
188        get_all,
189        frozen,
190        module = "qcs_api_client_common._qcs_api_client_common.configuration",
191        from_py_object
192    )
193)]
194/// The access (Bearer) and refresh (if available) tokens issued by an auth server.
195///
196/// These could be issued through one of multiple kinds of OAuth flows,
197/// see [`OAuthGrant::InteractiveLogin`] for implementation details.
198pub struct AuthTokens {
199    /// The access token.
200    pub access_token: SecretAccessToken,
201    /// The refresh token, if available.
202    pub refresh_token: Option<RefreshToken>,
203}
204
205/// Errors that can occur when attempting to perform an interactive login.
206#[derive(Debug, thiserror::Error)]
207pub enum LoginError {
208    /// Error that occurred while performing the Authorization Code (PKCE) flow.
209    #[error(transparent)]
210    Pkce(#[from] PkceLoginError),
211    /// Error that occurred while fetching the discovery document from the `OAuth2` issuer.
212    #[error(transparent)]
213    Discovery(#[from] DiscoveryError),
214    /// Error that occurred while making http requests.
215    #[error(transparent)]
216    Request(#[from] qcs_dependencies_client::reqwest::Error),
217}
218
219impl AuthTokens {
220    /// Performs an interactive login, returning the tokens the auth server issues.
221    ///
222    /// # Errors
223    ///
224    /// See [`LoginError`]
225    pub async fn interactive_login(
226        cancel_token: CancellationToken,
227        auth_server: &AuthServer,
228    ) -> Result<Self, LoginError> {
229        Self::interactive_login_with_redirect(cancel_token, auth_server, RedirectBinding::default())
230            .await
231    }
232
233    /// See [`Self::interactive_login`].
234    pub(crate) async fn interactive_login_with_redirect(
235        cancel_token: CancellationToken,
236        auth_server: &AuthServer,
237        redirect: RedirectBinding,
238    ) -> Result<Self, LoginError> {
239        let issuer = auth_server.issuer.clone();
240
241        let client = default_http_client()?;
242        let discovery = oidc::fetch_discovery(&client, &issuer).await?;
243
244        let response = pkce_login(
245            cancel_token,
246            PkceLoginRequest {
247                client_id: auth_server.client_id.clone(),
248                redirect,
249                discovery,
250                scopes: auth_server.scopes.clone(),
251            },
252        )
253        .await?;
254
255        Ok(Self {
256            access_token: SecretAccessToken::from(response.access_token().secret().clone()),
257            refresh_token: response
258                .refresh_token()
259                .map(|rt| RefreshToken::new(SecretRefreshToken::from(rt.secret().clone()))),
260        })
261    }
262
263    /// Returns the access token if it is valid, otherwise requests a new access token using the refresh token if available.
264    ///
265    /// # Errors
266    ///
267    /// See [`TokenError`]
268    pub async fn request_access_token(
269        &mut self,
270        auth_server: &AuthServer,
271    ) -> Result<SecretAccessToken, TokenError> {
272        if insecure_validate_token_exp(&self.access_token).is_ok() {
273            return Ok(self.access_token.clone());
274        }
275
276        if let Some(refresh_token) = &mut self.refresh_token {
277            let access_token = refresh_token.request_access_token(auth_server).await?;
278            self.access_token.clone_from(&access_token);
279            return Ok(access_token);
280        }
281
282        Err(TokenError::NoRefreshToken)
283    }
284}
285
286impl From<AuthTokens> for Credential {
287    fn from(value: AuthTokens) -> Self {
288        let mut token_payload = TokenPayload::default();
289        token_payload.access_token = Some(value.access_token);
290        token_payload.refresh_token = value.refresh_token.map(|rt| rt.refresh_token);
291
292        Self::TokenPayload(token_payload)
293    }
294}
295
296#[derive(Clone)]
297#[cfg_attr(feature = "python", derive(pyo3::FromPyObject, pyo3::IntoPyObject))]
298/// Specifies the [OAuth2 grant type](https://oauth.net/2/grant-types/) to use, along with the data
299/// needed to request said grant type.
300pub enum OAuthGrant {
301    /// Credentials that can be used to use with the [Refresh Token grant type](https://oauth.net/2/grant-types/refresh-token/).
302    RefreshToken(RefreshToken),
303    /// Payload that can be used to use the [Client Credentials grant type](https://oauth.net/2/grant-types/client-credentials/).
304    ClientCredentials(ClientCredentials),
305    /// Defers to a user provided function for access token requests.
306    ExternallyManaged(ExternallyManaged),
307    /// The tokens returned by an interactive login, i.e. one with a human in the loop.
308    ///
309    /// Currently that means the [Authorization Code grant with PKCE](https://oauth.net/2/pkce/).
310    InteractiveLogin(AuthTokens),
311}
312
313impl From<ExternallyManaged> for OAuthGrant {
314    fn from(v: ExternallyManaged) -> Self {
315        Self::ExternallyManaged(v)
316    }
317}
318
319impl From<ClientCredentials> for OAuthGrant {
320    fn from(v: ClientCredentials) -> Self {
321        Self::ClientCredentials(v)
322    }
323}
324
325impl From<RefreshToken> for OAuthGrant {
326    fn from(v: RefreshToken) -> Self {
327        Self::RefreshToken(v)
328    }
329}
330
331impl From<AuthTokens> for OAuthGrant {
332    fn from(v: AuthTokens) -> Self {
333        Self::InteractiveLogin(v)
334    }
335}
336
337impl OAuthGrant {
338    /// Request a new access token from the given issuer using this grant type and payload.
339    async fn request_access_token(
340        &mut self,
341        auth_server: &AuthServer,
342    ) -> Result<SecretAccessToken, TokenError> {
343        match self {
344            Self::RefreshToken(tokens) => tokens.request_access_token(auth_server).await,
345            Self::ClientCredentials(tokens) => tokens.request_access_token(auth_server).await,
346            Self::ExternallyManaged(tokens) => tokens
347                .request_access_token(auth_server)
348                .await
349                .map_err(|e| TokenError::ExternallyManaged(e.to_string())),
350            Self::InteractiveLogin(tokens) => tokens.request_access_token(auth_server).await,
351        }
352    }
353}
354
355impl std::fmt::Debug for OAuthGrant {
356    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
357        match self {
358            Self::RefreshToken(_) => f.write_str("RefreshToken"),
359            Self::ClientCredentials(_) => f.write_str("ClientCredentials"),
360            Self::ExternallyManaged(_) => f.write_str("ExternallyManaged"),
361            Self::InteractiveLogin(_) => f.write_str("InteractiveLogin"),
362        }
363    }
364}
365
366/// Manages the `OAuth2` authorization process and token lifecycle for accessing the QCS API.
367///
368/// This struct encapsulates the necessary information to request an access token
369/// from an authorization server, including the `OAuth2` grant type and any associated
370/// credentials or payload data.
371///
372/// # Fields
373///
374/// * `payload` - The `OAuth2` grant type and associated data that will be used to request an access token.
375/// * `access_token` - The access token currently in use, if any. If no token has been provided or requested yet, this will be `None`.
376/// * `auth_server` - The authorization server responsible for issuing tokens.
377#[derive(Clone)]
378#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
379#[cfg_attr(
380    feature = "python",
381    pyo3::pyclass(
382        module = "qcs_api_client_common._qcs_api_client_common.configuration",
383        frozen,
384        get_all,
385        from_py_object
386    )
387)]
388pub struct OAuthSession {
389    /// The grant type to use to request an access token.
390    payload: OAuthGrant,
391    /// The access token that is currently in use. None if no token has been requested yet.
392    access_token: Option<SecretAccessToken>,
393    /// The [`AuthServer`] that issues the tokens.
394    auth_server: AuthServer,
395}
396
397impl OAuthSession {
398    /// Initialize a new set of [`Credentials`] using a [`GrantPayload`].
399    ///
400    /// Optionally include an `access_token`, if not included, then one can be requested
401    /// with [`Self::request_access_token`].
402    #[must_use]
403    pub const fn new(
404        payload: OAuthGrant,
405        auth_server: AuthServer,
406        access_token: Option<SecretAccessToken>,
407    ) -> Self {
408        Self {
409            payload,
410            access_token,
411            auth_server,
412        }
413    }
414
415    /// Initialize a new set of [`Credentials`] using an [`ExternallyManaged`].
416    ///
417    /// Optionally include an `access_token`, if not included, then one can be requested
418    /// with [`Self::request_access_token`].
419    #[must_use]
420    pub const fn from_externally_managed(
421        tokens: ExternallyManaged,
422        auth_server: AuthServer,
423        access_token: Option<SecretAccessToken>,
424    ) -> Self {
425        Self::new(
426            OAuthGrant::ExternallyManaged(tokens),
427            auth_server,
428            access_token,
429        )
430    }
431
432    /// Initialize a new set of [`Credentials`] using a [`RefreshToken`].
433    ///
434    /// Optionally include an `access_token`, if not included, then one can be requested
435    /// with [`Self::request_access_token`].
436    #[must_use]
437    pub const fn from_refresh_token(
438        tokens: RefreshToken,
439        auth_server: AuthServer,
440        access_token: Option<SecretAccessToken>,
441    ) -> Self {
442        Self::new(OAuthGrant::RefreshToken(tokens), auth_server, access_token)
443    }
444
445    /// Initialize a new set of [`Credentials`] using [`ClientCredentials`].
446    ///
447    /// Optionally include an `access_token`, if not included, then one can be requested
448    /// with [`Self::request_access_token`].
449    #[must_use]
450    pub const fn from_client_credentials(
451        tokens: ClientCredentials,
452        auth_server: AuthServer,
453        access_token: Option<SecretAccessToken>,
454    ) -> Self {
455        Self::new(
456            OAuthGrant::ClientCredentials(tokens),
457            auth_server,
458            access_token,
459        )
460    }
461
462    /// Initialize a new set of [`Credentials`] using [`AuthTokens`].
463    ///
464    /// Optionally include an `access_token`, if not included, then one can be requested
465    /// with [`Self::request_access_token`].
466    #[must_use]
467    pub const fn from_interactive_login(
468        tokens: AuthTokens,
469        auth_server: AuthServer,
470        access_token: Option<SecretAccessToken>,
471    ) -> Self {
472        Self::new(
473            OAuthGrant::InteractiveLogin(tokens),
474            auth_server,
475            access_token,
476        )
477    }
478
479    /// Get the current access token.
480    ///
481    /// This is an unvalidated copy of the access token. Meaning it can become stale, or may
482    /// even be already be stale. See [`Self::validate`] and [`Self::request_access_token`].
483    ///
484    /// # Errors
485    ///
486    /// - [`TokenError::NoAccessToken`] if there is no access token
487    pub fn access_token(&self) -> Result<&SecretAccessToken, TokenError> {
488        self.access_token.as_ref().ok_or(TokenError::NoAccessToken)
489    }
490
491    /// Get the payload used to request an access token.
492    #[must_use]
493    pub const fn payload(&self) -> &OAuthGrant {
494        &self.payload
495    }
496
497    /// Request and return an updated access token using these credentials.
498    ///
499    /// # Errors
500    ///
501    /// See [`TokenError`]
502    #[allow(clippy::missing_panics_doc)]
503    pub async fn request_access_token(&mut self) -> Result<&SecretAccessToken, TokenError> {
504        let access_token = self.payload.request_access_token(&self.auth_server).await?;
505        Ok(self.access_token.insert(access_token))
506    }
507
508    /// The [`AuthServer`] that issues the tokens.
509    #[must_use]
510    pub const fn auth_server(&self) -> &AuthServer {
511        &self.auth_server
512    }
513
514    /// Validate the access token, returning it if it is valid, or an error describing why it is
515    /// invalid.
516    ///
517    /// # Errors
518    ///
519    /// - [`TokenError::NoAccessToken`] if an access token has not been requested.
520    /// - [`TokenError::InvalidAccessToken`] if the access token is invalid.
521    pub fn validate(&self) -> Result<SecretAccessToken, TokenError> {
522        let access_token = self.access_token()?;
523        insecure_validate_token_exp(access_token)?;
524        Ok(access_token.clone())
525    }
526}
527
528/// Validates the access token's format and `exp` claim, but no other claims or
529/// signature. We do this only to determine if the token is expired and needs refreshing,
530/// there is no way to securely validate the token's signature on the client side.
531pub(crate) fn insecure_validate_token_exp(
532    access_token: &SecretAccessToken,
533) -> Result<(), TokenError> {
534    let placeholder_key = DecodingKey::from_secret(&[]);
535    let mut validation = Validation::new(Algorithm::RS256);
536    validation.validate_exp = true;
537    validation.leeway = 60;
538    validation.validate_aud = false;
539    validation.insecure_disable_signature_validation();
540
541    jsonwebtoken::decode::<toml::Value>(access_token.secret(), &placeholder_key, &validation)
542        .map(|_| ())
543        .map_err(TokenError::InvalidAccessToken)
544}
545
546impl std::fmt::Debug for OAuthSession {
547    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
548        let token_populated = if self.access_token.is_some() {
549            Some(())
550        } else {
551            None
552        };
553        f.debug_struct("OAuthSession")
554            .field("payload", &self.payload)
555            .field("access_token", &token_populated)
556            .field("auth_server", &self.auth_server)
557            .finish()
558    }
559}
560
561/// Persists `oauth_session`'s tokens to the secrets file backing `source`, if any.
562///
563/// This is a no-op if `source` is not file-backed ([`ConfigSource::Default`] or
564/// [`ConfigSource::Builder`]), or if the secrets file is read-only (see [`Secrets::is_read_only`]).
565///
566/// Every code path that obtains a new or refreshed [`OAuthSession`] (whether through the
567/// [`TokenDispatcher`], or through [`ClientConfiguration::load_with_login`]'s manual refresh and
568/// interactive login branches) should call this so that a rotated refresh token isn't silently
569/// dropped.
570/// Otherwise, the next process to load the profile will retry a stale, already-consumed refresh
571/// token and be forced back into an interactive login.
572///
573/// # Errors
574///
575/// See [`WriteError`]
576pub(crate) async fn persist_oauth_session(
577    oauth_session: &OAuthSession,
578    source: &ConfigSource,
579    credentials_name: &str,
580) -> Result<(), WriteError> {
581    let ConfigSource::File {
582        settings_path: _,
583        secrets_path,
584    } = source
585    else {
586        return Ok(());
587    };
588
589    // Persist the fresh refresh token if the grant carries one, so that a rotated
590    // refresh token isn't lost on the next load. Both the interactive-login and
591    // refresh-token grants can hold a refresh token that the auth server may have rotated.
592    let refresh_token = match &oauth_session.payload {
593        OAuthGrant::InteractiveLogin(payload) => {
594            payload.refresh_token.as_ref().map(|rt| &rt.refresh_token)
595        }
596        OAuthGrant::RefreshToken(payload) => Some(&payload.refresh_token),
597        OAuthGrant::ExternallyManaged(_) | OAuthGrant::ClientCredentials(_) => return Ok(()),
598    };
599
600    if Secrets::is_read_only(secrets_path).await? {
601        #[cfg(feature = "tracing")]
602        tracing::debug!(
603            "Skipping write of refreshed tokens to read-only secrets file: {:?}",
604            secrets_path
605        );
606        return Ok(());
607    }
608
609    // Nothing to persist without an access token; this shouldn't happen for a session that was
610    // just successfully refreshed or logged in, but there's nothing useful to write otherwise.
611    let Ok(access_token) = oauth_session.access_token() else {
612        return Ok(());
613    };
614
615    let now = OffsetDateTime::now_utc();
616    Secrets::write_tokens(
617        secrets_path,
618        credentials_name,
619        refresh_token,
620        access_token,
621        now,
622    )
623    .await
624}
625
626/// A wrapper for [`OAuthSession`] that provides thread-safe access to the inner tokens.
627#[derive(Clone, Debug)]
628#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
629#[cfg_attr(
630    feature = "python",
631    pyo3::pyclass(
632        module = "qcs_api_client_common._qcs_api_client_common.configuration",
633        frozen,
634        from_py_object
635    )
636)]
637pub struct TokenDispatcher {
638    lock: Arc<RwLock<OAuthSession>>,
639    refreshing: Arc<Mutex<bool>>,
640    notify_refreshed: Arc<Notify>,
641}
642
643impl From<OAuthSession> for TokenDispatcher {
644    fn from(value: OAuthSession) -> Self {
645        Self {
646            lock: Arc::new(RwLock::new(value)),
647            refreshing: Arc::new(Mutex::new(false)),
648            notify_refreshed: Arc::new(Notify::new()),
649        }
650    }
651}
652
653impl TokenDispatcher {
654    /// Executes a user-provided closure on a reference to the `Tokens` instance managed by the
655    /// dispatcher.
656    ///
657    /// This function locks the mutex, safely exposing the protected `Tokens` instance to the provided closure `f`.
658    /// It is designed to allow safe and controlled access to the `Tokens` instance for reading its state.
659    ///
660    /// # Parameters
661    /// - `f`: A closure that takes a reference to `Tokens` and returns a value of type `O`. The closure is called
662    ///   with the `Tokens` instance as an argument once the mutex is successfully locked.
663    pub async fn use_tokens<F, O>(&self, f: F) -> O
664    where
665        F: FnOnce(&OAuthSession) -> O + Send,
666    {
667        let tokens = self.lock.read().await;
668        f(&tokens)
669    }
670
671    /// Get a copy of the current access token.
672    #[must_use]
673    pub async fn tokens(&self) -> OAuthSession {
674        self.use_tokens(Clone::clone).await
675    }
676
677    /// Refreshes the tokens. Readers will be blocked until the refresh is complete.
678    ///
679    /// # Errors
680    ///
681    /// See [`TokenError`]
682    pub async fn refresh(
683        &self,
684        source: &ConfigSource,
685        credentials_name: &str,
686    ) -> Result<OAuthSession, TokenError> {
687        self.managed_refresh(Self::perform_refresh, source, credentials_name)
688            .await
689    }
690
691    /// Validate the access token, returning it if it is valid, or an error describing why it is
692    /// invalid.
693    ///
694    /// # Errors
695    ///
696    /// - [`TokenError::NoAccessToken`] if there is no access token
697    /// - [`TokenError::InvalidAccessToken`] if the access token is invalid
698    pub async fn validate(&self) -> Result<SecretAccessToken, TokenError> {
699        self.use_tokens(OAuthSession::validate).await
700    }
701
702    /// If tokens are already being refreshed, wait and return the updated tokens. Otherwise, run
703    /// ``refresh_fn``.
704    async fn managed_refresh<F, Fut>(
705        &self,
706        refresh_fn: F,
707        source: &ConfigSource,
708        credentials_name: &str,
709    ) -> Result<OAuthSession, TokenError>
710    where
711        F: FnOnce(Arc<RwLock<OAuthSession>>) -> Fut + Send,
712        Fut: Future<Output = Result<OAuthSession, TokenError>> + Send,
713    {
714        let mut is_refreshing = self.refreshing.lock().await;
715
716        if *is_refreshing {
717            drop(is_refreshing);
718            self.notify_refreshed.notified().await;
719            return Ok(self.tokens().await);
720        }
721
722        *is_refreshing = true;
723        drop(is_refreshing);
724
725        let oauth_session = refresh_fn(self.lock.clone()).await?;
726
727        let write_result = persist_oauth_session(&oauth_session, source, credentials_name).await;
728
729        // Always clean up the refreshing lock, even if write failed
730        *self.refreshing.lock().await = false;
731        self.notify_refreshed.notify_waiters();
732
733        // If write failed, return error with the valid oauth_session
734        if let Err(error) = write_result {
735            return Err(TokenError::Write {
736                error,
737                oauth_session: Box::new(oauth_session),
738            });
739        }
740
741        Ok(oauth_session)
742    }
743
744    /// Refreshes the tokens. Readers will be blocked until the refresh is complete. Returns a copy
745    /// of the updated [`Credentials`]
746    ///
747    /// # Errors
748    ///
749    /// See [`TokenError`]
750    async fn perform_refresh(lock: Arc<RwLock<OAuthSession>>) -> Result<OAuthSession, TokenError> {
751        let mut credentials = lock.write().await;
752        credentials.request_access_token().await?;
753        Ok(credentials.clone())
754    }
755}
756
757pub(crate) type RefreshResult =
758    Pin<Box<dyn Future<Output = Result<String, Box<dyn std::error::Error + Send + Sync>>> + Send>>;
759
760/// A function that asynchronously refreshes a token.
761pub type RefreshFunction = Box<dyn (Fn(AuthServer) -> RefreshResult) + Send + Sync>;
762
763/// A struct that manages access tokens by utilizing a user-provided refresh function.
764///
765/// The [`ExternallyManaged`] struct allows users to define custom logic for
766/// fetching or refreshing access tokens.
767#[derive(Clone)]
768#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
769#[cfg_attr(
770    feature = "python",
771    pyo3::pyclass(
772        module = "qcs_api_client_common._qcs_api_client_common.configuration",
773        frozen,
774        from_py_object
775    )
776)]
777pub struct ExternallyManaged {
778    refresh_function: Arc<RefreshFunction>,
779}
780
781impl ExternallyManaged {
782    /// Creates a new [`ExternallyManaged`] instance from a [`RefreshFunction`].
783    ///
784    /// Consider using [`ExternallyManaged::from_async`], and [`ExternallyManaged::from_sync`], if
785    /// they better fit your use case.
786    ///
787    /// # Arguments
788    ///
789    /// * `refresh_function` - A function or closure that asynchronously refreshes a token.
790    ///
791    /// # Example
792    ///
793    /// ```
794    /// use qcs_api_client_common::configuration::{settings::AuthServer, tokens::ExternallyManaged, TokenError};
795    /// use std::future::Future;
796    /// use std::pin::Pin;
797    /// use std::boxed::Box;
798    /// use std::error::Error;
799    ///
800    /// async fn example_refresh_function(_auth_server: AuthServer) -> Result<String, Box<dyn Error
801    /// + Send + Sync>> {
802    ///     Ok("new_token_value".to_string())
803    /// }
804    /// let token_manager = ExternallyManaged::new(|auth_server| Box::pin(example_refresh_function(auth_server)));
805    /// ```
806    pub fn new(
807        refresh_function: impl Fn(AuthServer) -> RefreshResult + Send + Sync + 'static,
808    ) -> Self {
809        Self {
810            refresh_function: Arc::new(Box::new(refresh_function)),
811        }
812    }
813
814    /// Constructs a new [`ExternallyManaged`] instance using an async function or closure.
815    ///
816    /// This method simplifies the creation of the [`ExternallyManaged`] instance by handling
817    /// the boxing and pinning of the future internally.
818    ///
819    /// # Arguments
820    ///
821    /// * `refresh_function` - An async function or closure that returns a [`Future`] which, when awaited,
822    ///   produces a [`Result<String, TokenError>`].
823    ///
824    /// # Example
825    ///
826    /// ```
827    /// use qcs_api_client_common::configuration::{settings::AuthServer, tokens::ExternallyManaged, TokenError};
828    /// use tokio::runtime::Runtime;
829    /// use std::error::Error;
830    ///
831    /// async fn example_refresh_function(_auth_server: AuthServer) -> Result<String, Box<dyn Error
832    /// + Send + Sync>> {
833    ///     Ok("new_token_value".to_string())
834    /// }
835    ///
836    /// let token_manager = ExternallyManaged::from_async(example_refresh_function);
837    ///
838    /// let rt = Runtime::new().unwrap();
839    /// rt.block_on(async {
840    ///     match token_manager.request_access_token(&AuthServer::default()).await {
841    ///         Ok(token) => println!("Token: {token:?}"),
842    ///         Err(e) => println!("Failed to refresh token: {:?}", e),
843    ///     }
844    /// });
845    /// ```
846    pub fn from_async<F, Fut>(refresh_function: F) -> Self
847    where
848        F: Fn(AuthServer) -> Fut + Send + Sync + 'static,
849        Fut: Future<Output = Result<String, Box<dyn std::error::Error + Send + Sync>>>
850            + Send
851            + 'static,
852    {
853        Self {
854            refresh_function: Arc::new(Box::new(move |auth_server| {
855                Box::pin(refresh_function(auth_server))
856            })),
857        }
858    }
859
860    /// Constructs a new [`ExternallyManaged`] instance using a synchronous function.
861    ///
862    /// The synchronous function is wrapped in an async block to fit the expected signature.
863    ///
864    /// # Arguments
865    ///
866    /// * `refresh_function` - A synchronous function that returns a [`Result<String, TokenError>`].
867    ///
868    /// # Example
869    ///
870    /// ```
871    /// use qcs_api_client_common::configuration::{settings::AuthServer, tokens::ExternallyManaged, TokenError};
872    /// use tokio::runtime::Runtime;
873    /// use std::error::Error;
874    ///
875    /// fn example_sync_refresh_function(_auth_server: AuthServer) -> Result<String, Box<dyn Error
876    /// + Send + Sync>> {
877    ///     Ok("sync_token_value".to_string())
878    /// }
879    ///
880    /// let token_manager = ExternallyManaged::from_sync(example_sync_refresh_function);
881    ///
882    /// let rt = Runtime::new().unwrap();
883    /// rt.block_on(async {
884    ///     match token_manager.request_access_token(&AuthServer::default()).await {
885    ///         Ok(token) => println!("Token: {token:?}"),
886    ///         Err(e) => println!("Failed to refresh token: {:?}", e),
887    ///     }
888    /// });
889    /// ```
890    pub fn from_sync(
891        refresh_function: impl Fn(
892            AuthServer,
893        ) -> Result<String, Box<dyn std::error::Error + Send + Sync>>
894        + Send
895        + Sync
896        + 'static,
897    ) -> Self {
898        Self {
899            refresh_function: Arc::new(Box::new(move |auth_server| {
900                let result = refresh_function(auth_server);
901                Box::pin(async move { result })
902            })),
903        }
904    }
905
906    /// Request an updated access token using the provided refresh function.
907    ///
908    /// # Errors
909    ///
910    /// Errors are propagated from the refresh function.
911    pub async fn request_access_token(
912        &self,
913        auth_server: &AuthServer,
914    ) -> Result<SecretAccessToken, Box<dyn std::error::Error + Send + Sync>> {
915        (self.refresh_function)(auth_server.clone())
916            .await
917            .map(SecretAccessToken::from)
918    }
919}
920
921impl std::fmt::Debug for ExternallyManaged {
922    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
923        f.debug_struct("ExternallyManaged")
924            .field(
925                "refresh_function",
926                &"Fn() -> Pin<Box<dyn Future<Output = Result<String, TokenError>> + Send>>",
927            )
928            .finish()
929    }
930}
931
932#[derive(Debug, Serialize, Deserialize)]
933pub(super) struct TokenRefreshRequest<'a> {
934    grant_type: &'static str,
935    client_id: &'a str,
936    refresh_token: &'a str,
937}
938
939impl<'a> TokenRefreshRequest<'a> {
940    pub(super) const fn new(client_id: &'a str, refresh_token: &'a str) -> Self {
941        Self {
942            grant_type: "refresh_token",
943            client_id,
944            refresh_token,
945        }
946    }
947}
948
949#[derive(Debug, Serialize, Deserialize)]
950pub(super) struct ClientCredentialsRequest {
951    grant_type: &'static str,
952    scope: Option<&'static str>,
953}
954
955impl ClientCredentialsRequest {
956    pub(super) const fn new(scope: Option<&'static str>) -> Self {
957        Self {
958            grant_type: "client_credentials",
959            scope,
960        }
961    }
962}
963
964#[derive(Deserialize, Debug, Serialize)]
965pub(super) struct RefreshTokenResponse {
966    pub(super) refresh_token: Option<SecretRefreshToken>,
967    pub(super) access_token: SecretAccessToken,
968}
969
970/// Get and refresh access tokens
971#[async_trait::async_trait]
972pub trait TokenRefresher: Clone + std::fmt::Debug + Send {
973    /// The type to be returned in the event of a error during getting or
974    /// refreshing an access token
975    type Error;
976
977    /// Get and validate the current access token, refreshing it if it doesn't exist or is invalid.
978    async fn validated_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
979
980    /// Get the current access token, if any
981    async fn get_access_token(&self) -> Result<Option<SecretAccessToken>, Self::Error>;
982
983    /// Get a fresh access token
984    async fn refresh_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
985
986    /// Get the base URL for requests
987    #[cfg(feature = "tracing")]
988    fn base_url(&self) -> &str;
989
990    /// Get the tracing configuration
991    #[cfg(feature = "tracing-config")]
992    fn tracing_configuration(&self) -> Option<&TracingConfiguration>;
993
994    /// Returns whether the given URL should be traced. Following
995    /// [`TracingConfiguration::is_enabled`], this defaults to `true`.
996    #[cfg(feature = "tracing")]
997    #[allow(clippy::needless_return)]
998    fn should_trace(&self, url: &UrlPatternMatchInput) -> bool {
999        #[cfg(not(feature = "tracing-config"))]
1000        {
1001            let _ = url;
1002            return true;
1003        }
1004
1005        #[cfg(feature = "tracing-config")]
1006        self.tracing_configuration()
1007            .is_none_or(|config| config.is_enabled(url))
1008    }
1009}
1010
1011#[async_trait::async_trait]
1012impl TokenRefresher for ClientConfiguration {
1013    type Error = TokenError;
1014
1015    async fn validated_access_token(&self) -> Result<SecretAccessToken, Self::Error> {
1016        self.get_bearer_access_token().await
1017    }
1018
1019    async fn refresh_access_token(&self) -> Result<SecretAccessToken, Self::Error> {
1020        match self.refresh().await {
1021            Ok(session) => Ok(session.access_token()?.clone()),
1022            Err(TokenError::Write {
1023                error: _error,
1024                oauth_session,
1025            }) => {
1026                // Token refresh succeeded but persistence failed. Extract and return the access token from the error.
1027                #[cfg(feature = "tracing")]
1028                tracing::warn!(
1029                    "Token refresh succeeded but failed to persist: {_error}. Returning access token from error.",
1030                );
1031                Ok(oauth_session.access_token()?.clone())
1032            }
1033            Err(e) => Err(e),
1034        }
1035    }
1036
1037    async fn get_access_token(&self) -> Result<Option<SecretAccessToken>, Self::Error> {
1038        Ok(Some(self.oauth_session().await?.access_token()?.clone()))
1039    }
1040
1041    #[cfg(feature = "tracing")]
1042    fn base_url(&self) -> &str {
1043        &self.grpc_api_url
1044    }
1045
1046    #[cfg(feature = "tracing-config")]
1047    fn tracing_configuration(&self) -> Option<&TracingConfiguration> {
1048        self.tracing_configuration.as_ref()
1049    }
1050}
1051
1052/// Get a default http client.
1053///
1054/// # Errors
1055///
1056/// Returns an error if the underlying `reqwest` client fails to build.
1057pub fn default_http_client()
1058-> Result<qcs_dependencies_client::reqwest::Client, qcs_dependencies_client::reqwest::Error> {
1059    qcs_dependencies_client::reqwest::Client::builder()
1060        .timeout(std::time::Duration::from_secs(10))
1061        .build()
1062}
1063
1064#[cfg(test)]
1065mod test {
1066    #![allow(clippy::result_large_err, reason = "happens in figment tests")]
1067
1068    use std::time::Duration;
1069
1070    use super::*;
1071    use httpmock::prelude::*;
1072    use rstest::rstest;
1073    use time::format_description::well_known::Rfc3339;
1074    use tokio::time::Instant;
1075    use toml_edit::DocumentMut;
1076
1077    #[tokio::test]
1078    async fn test_tokens_blocked_during_refresh() {
1079        let mock_server = MockServer::start_async().await;
1080
1081        let oidc_mock = mock_server
1082            .mock_async(|when, then| {
1083                when.method(GET).path("/.well-known/openid-configuration");
1084                then.status(200)
1085                    .json_body_obj(&oidc::Discovery::new_for_test(
1086                        mock_server.base_url().parse().unwrap(),
1087                    ));
1088            })
1089            .await;
1090
1091        let issuer_mock = mock_server
1092            .mock_async(|when, then| {
1093                when.method(POST).path("/v1/token");
1094
1095                then.status(200)
1096                    .delay(Duration::from_secs(3))
1097                    .json_body_obj(&RefreshTokenResponse {
1098                        access_token: SecretAccessToken::from("new_access"),
1099                        refresh_token: Some(SecretRefreshToken::from("new_refresh")),
1100                    });
1101            })
1102            .await;
1103
1104        let original_tokens = OAuthSession::from_refresh_token(
1105            RefreshToken::new(SecretRefreshToken::from("refresh")),
1106            AuthServer {
1107                client_id: "client_id".to_string(),
1108                issuer: mock_server.base_url(),
1109                scopes: None,
1110            },
1111            None,
1112        );
1113        let dispatcher: TokenDispatcher = original_tokens.clone().into();
1114        let dispatcher_clone1 = dispatcher.clone();
1115        let dispatcher_clone2 = dispatcher.clone();
1116
1117        let refresh_duration = Duration::from_secs(3);
1118
1119        let start_write = Instant::now();
1120        let write_future = tokio::spawn(async move {
1121            dispatcher_clone1
1122                .refresh(&ConfigSource::Default, "")
1123                .await
1124                .unwrap()
1125        });
1126
1127        let start_read = Instant::now();
1128        let read_future = tokio::spawn(async move { dispatcher_clone2.tokens().await });
1129
1130        let _ = write_future.await.unwrap();
1131        let read_result = read_future.await.unwrap();
1132
1133        let write_duration = start_write.elapsed();
1134        let read_duration = start_read.elapsed();
1135
1136        oidc_mock.assert_async().await;
1137        issuer_mock.assert_async().await;
1138
1139        assert!(
1140            write_duration >= refresh_duration,
1141            "Write operation did not take enough time"
1142        );
1143        assert!(
1144            read_duration >= refresh_duration,
1145            "Read operation was not blocked by the write operation"
1146        );
1147        assert_eq!(
1148            read_result.access_token.unwrap(),
1149            SecretAccessToken::from("new_access")
1150        );
1151        if let OAuthGrant::RefreshToken(payload) = read_result.payload {
1152            assert_eq!(
1153                payload.refresh_token,
1154                SecretRefreshToken::from("new_refresh")
1155            );
1156        } else {
1157            panic!(
1158                "Expected RefreshToken payload, got {:?}",
1159                read_result.payload
1160            );
1161        }
1162    }
1163
1164    /// When the auth server rejects a refresh token request (e.g. the refresh token was revoked
1165    /// or has expired), the failure should still surface as a normal error - not panic - even
1166    /// though the response body is read for logging before the error is returned.
1167    #[tokio::test]
1168    async fn test_refresh_token_request_rejected_by_auth_server() {
1169        let mock_server = MockServer::start_async().await;
1170
1171        let oidc_mock = mock_server
1172            .mock_async(|when, then| {
1173                when.method(GET).path("/.well-known/openid-configuration");
1174                then.status(200)
1175                    .json_body_obj(&oidc::Discovery::new_for_test(
1176                        mock_server.base_url().parse().unwrap(),
1177                    ));
1178            })
1179            .await;
1180
1181        let issuer_mock = mock_server
1182            .mock_async(|when, then| {
1183                when.method(POST).path("/v1/token");
1184                then.status(400).json_body_obj(&serde_json::json!({
1185                    "error": "invalid_grant",
1186                    "error_description": "Unknown or invalid refresh token.",
1187                }));
1188            })
1189            .await;
1190
1191        let mut refresh_token = RefreshToken::new(SecretRefreshToken::from("revoked_refresh"));
1192        let auth_server = AuthServer {
1193            client_id: "client_id".to_string(),
1194            issuer: mock_server.base_url(),
1195            scopes: None,
1196        };
1197
1198        let result = refresh_token.request_access_token(&auth_server).await;
1199
1200        oidc_mock.assert_async().await;
1201        issuer_mock.assert_async().await;
1202
1203        assert!(
1204            result.is_err(),
1205            "a rejected refresh token request should be an error, got {result:?}"
1206        );
1207    }
1208
1209    #[rstest]
1210    fn test_qcs_secrets_readonly(
1211        #[values(
1212            (Some("TRUE"), true),
1213            (Some("tRue"), true),
1214            (Some("true"), true),
1215            (Some("YES"), true),
1216            (Some("yEs"), true),
1217            (Some("yes"), true),
1218            (Some("1"), true),
1219            (Some("2"), false),
1220            (Some("other"), false),
1221            (Some(""), false),
1222            (None, false),
1223        )]
1224        read_only_values: (Option<&str>, bool),
1225        #[values(true, false)] read_only_perm: bool,
1226    ) {
1227        let (maybe_read_only_env, env_is_read_only) = read_only_values;
1228        let expected_update = !env_is_read_only && !read_only_perm;
1229        figment::Jail::expect_with(|jail| {
1230            let profile_name = "test";
1231            let initial_access_token = "initial_access_token";
1232            let initial_refresh_token = "initial_refresh_token";
1233
1234            let initial_secrets_file_contents = format!(
1235                r#"
1236[credentials]
1237[credentials.{profile_name}]
1238[credentials.{profile_name}.token_payload]
1239access_token = "{initial_access_token}"
1240expires_in = 3600
1241id_token = "id_token"
1242refresh_token = "{initial_refresh_token}"
1243scope = "offline_access openid profile email"
1244token_type = "Bearer"
1245updated_at = "2024-01-01T00:00:00Z"
1246"#
1247            );
1248
1249            // Ignore any existing environment variables.
1250            jail.clear_env();
1251
1252            // Create a temporary secrets file
1253            let secrets_path = "secrets.toml";
1254            jail.create_file(secrets_path, initial_secrets_file_contents.as_str())
1255                .expect("should create test secrets.toml");
1256
1257            if read_only_perm {
1258                let mut permissions = std::fs::metadata(secrets_path)
1259                    .expect("Should be able to get file metadata")
1260                    .permissions();
1261                permissions.set_readonly(true);
1262                std::fs::set_permissions(secrets_path, permissions)
1263                    .expect("Should be able to set file permissions");
1264            }
1265
1266            let rt = tokio::runtime::Runtime::new().unwrap();
1267            rt.block_on(async {
1268                let mock_server = MockServer::start_async().await;
1269
1270                let oidc_mock = mock_server
1271                    .mock_async(|when, then| {
1272                        when.method(GET).path("/.well-known/openid-configuration");
1273                        then.status(200)
1274                            .json_body_obj(&oidc::Discovery::new_for_test(mock_server.base_url().parse().unwrap()));
1275                    })
1276                    .await;
1277
1278                // Set up the mock token endpoint
1279                let new_access_token = SecretAccessToken::from("new_access_token");
1280                let issuer_mock = mock_server
1281                    .mock_async(|when, then| {
1282                        when.method(POST).path("/v1/token");
1283                        then.status(200).json_body_obj(&RefreshTokenResponse {
1284                            access_token: new_access_token.clone(),
1285                            refresh_token: Some(SecretRefreshToken::from(initial_refresh_token)),
1286                        });
1287                    })
1288                    .await;
1289
1290                // Create tokens and dispatcher
1291                let original_tokens = OAuthSession::from_refresh_token(
1292                    RefreshToken::new(SecretRefreshToken::from(initial_refresh_token)),
1293                    AuthServer { client_id: "client_id".to_string(), issuer: mock_server.base_url(), scopes: None },
1294                    Some(SecretAccessToken::from(initial_refresh_token)),
1295                );
1296                let dispatcher: TokenDispatcher = original_tokens.into();
1297
1298                // Test with QCS_SECRETS_READ_ONLY set first
1299                jail.set_env("QCS_SECRETS_FILE_PATH", "secrets.toml");
1300                jail.set_env("QCS_PROFILE_NAME", "test");
1301                if let Some(read_only_env) = maybe_read_only_env {
1302                    jail.set_env("QCS_SECRETS_READ_ONLY", read_only_env);
1303                }
1304
1305                let before_refresh = OffsetDateTime::now_utc();
1306
1307                dispatcher
1308                    .refresh(
1309                        &ConfigSource::File {
1310                            settings_path: "".into(),
1311                            secrets_path: "secrets.toml".into(),
1312                        },
1313                        profile_name,
1314                    )
1315                    .await
1316                    .unwrap();
1317
1318                oidc_mock.assert_async().await;
1319                issuer_mock.assert_async().await;
1320
1321                // Verify the file was not updated if QCS_SECRETS_READ_ONLY is set truthy
1322                let content = std::fs::read_to_string("secrets.toml").unwrap();
1323                if !expected_update {
1324                    assert!(
1325                        content.eq(initial_secrets_file_contents.as_str()),
1326                        "File should not be updated when QCS_SECRETS_READ_ONLY is set or file permissions are read-only"
1327                    );
1328                    return;
1329                }
1330
1331                // Verify the file was updated
1332                let mut toml = std::fs::read_to_string(secrets_path)
1333                    .unwrap()
1334                    .parse::<DocumentMut>()
1335                    .unwrap();
1336
1337                let token_payload = toml
1338                    .get_mut("credentials")
1339                    .and_then(|credentials| {
1340                        credentials.get_mut(profile_name)?.get_mut("token_payload")
1341                    })
1342                    .expect("Should be able to get token_payload table");
1343
1344                let access_token = token_payload.get("access_token").unwrap().as_str().map(str::to_string).map(SecretAccessToken::from);
1345
1346                assert_eq!(
1347                    access_token,
1348                    Some(new_access_token)
1349                );
1350
1351                assert!(
1352                    OffsetDateTime::parse(
1353                        token_payload.get("updated_at").unwrap().as_str().unwrap(),
1354                        &Rfc3339
1355                    )
1356                    .unwrap()
1357                        > before_refresh
1358                );
1359
1360                let content = std::fs::read_to_string("secrets.toml").unwrap();
1361                assert!(
1362                content.contains("new_access_token"),
1363                "File should be updated with new access token when QCS_SECRETS_READ_ONLY is not set or is set but disabled, and file permissions allow writing"
1364                );
1365            });
1366            Ok(())
1367        });
1368    }
1369
1370    /// When the auth server rotates the refresh token, a [`OAuthGrant::RefreshToken`] grant should
1371    /// persist the new refresh token to the secrets file (not just the access token).
1372    #[test]
1373    fn test_refresh_token_grant_persists_rotated_refresh_token() {
1374        let initial_refresh_token = "initial_refresh_token";
1375        let rotated_refresh_token = "rotated_refresh_token";
1376        let new_access_token = "new_access_token";
1377
1378        figment::Jail::expect_with(|jail| {
1379            jail.clear_env();
1380
1381            let secrets_path = "secrets.toml";
1382            let initial_secrets_file_contents = format!(
1383                r#"
1384[credentials]
1385[credentials.test]
1386[credentials.test.token_payload]
1387access_token = "initial_access_token"
1388refresh_token = "{initial_refresh_token}"
1389updated_at = "2024-01-01T00:00:00Z"
1390"#
1391            );
1392            jail.create_file(secrets_path, &initial_secrets_file_contents)
1393                .expect("should create test secrets.toml");
1394
1395            let rt = tokio::runtime::Runtime::new().unwrap();
1396            rt.block_on(async {
1397                let mock_server = MockServer::start_async().await;
1398                let oidc_mock = mock_server
1399                    .mock_async(|when, then| {
1400                        when.method(GET).path("/.well-known/openid-configuration");
1401                        then.status(200)
1402                            .json_body_obj(&oidc::Discovery::new_for_test(
1403                                mock_server.base_url().parse().unwrap(),
1404                            ));
1405                    })
1406                    .await;
1407                let issuer_mock = mock_server
1408                    .mock_async(|when, then| {
1409                        when.method(POST).path("/v1/token");
1410                        then.status(200).json_body_obj(&RefreshTokenResponse {
1411                            access_token: SecretAccessToken::from(new_access_token),
1412                            refresh_token: Some(SecretRefreshToken::from(rotated_refresh_token)),
1413                        });
1414                    })
1415                    .await;
1416
1417                let dispatcher: TokenDispatcher = OAuthSession::from_refresh_token(
1418                    RefreshToken::new(SecretRefreshToken::from(initial_refresh_token)),
1419                    AuthServer {
1420                        client_id: "client_id".to_string(),
1421                        issuer: mock_server.base_url(),
1422                        scopes: None,
1423                    },
1424                    Some(SecretAccessToken::from("initial_access_token")),
1425                )
1426                .into();
1427
1428                dispatcher
1429                    .refresh(
1430                        &ConfigSource::File {
1431                            settings_path: "".into(),
1432                            secrets_path: secrets_path.into(),
1433                        },
1434                        "test",
1435                    )
1436                    .await
1437                    .expect("refresh should succeed");
1438
1439                oidc_mock.assert_async().await;
1440                issuer_mock.assert_async().await;
1441            });
1442
1443            // The rotated refresh token (and the new access token) should be persisted.
1444            let Credential::TokenPayload(payload) = Secrets::load_from_path(&secrets_path.into())
1445                .expect("should load secrets")
1446                .credentials
1447                .remove("test")
1448                .expect("should have test credentials")
1449            else {
1450                panic!("expected a token payload credential");
1451            };
1452            assert_eq!(
1453                payload.refresh_token.unwrap(),
1454                SecretRefreshToken::from(rotated_refresh_token),
1455                "rotated refresh token should be persisted to the secrets file"
1456            );
1457            assert_eq!(
1458                payload.access_token.unwrap(),
1459                SecretAccessToken::from(new_access_token),
1460                "new access token should be persisted to the secrets file"
1461            );
1462
1463            Ok(())
1464        });
1465    }
1466
1467    #[test]
1468    fn test_auth_session_debug_fmt() {
1469        let session = OAuthSession {
1470            payload: OAuthGrant::ClientCredentials(ClientCredentials::new(
1471                "hidden_id",
1472                "hidden_secret",
1473            )),
1474            access_token: Some(SecretAccessToken::from("token")),
1475            auth_server: AuthServer {
1476                client_id: "some_id".into(),
1477                issuer: "some_url".into(),
1478                scopes: None,
1479            },
1480        };
1481
1482        assert_eq!(
1483            "OAuthSession { payload: ClientCredentials, access_token: Some(()), auth_server: AuthServer { client_id: \"some_id\", issuer: \"some_url\", scopes: None } }",
1484            &format!("{session:?}")
1485        );
1486    }
1487}