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