1use 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, 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#[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 pub refresh_token: SecretRefreshToken,
46}
47
48impl RefreshToken {
49 #[must_use]
51 pub const fn new(refresh_token: SecretRefreshToken) -> Self {
52 Self { refresh_token }
53 }
54
55 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 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#[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 pub client_id: String,
128 pub client_secret: ClientSecret,
130}
131
132impl ClientCredentials {
133 #[must_use]
134 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 #[must_use]
144 pub fn client_id(&self) -> &str {
145 &self.client_id
146 }
147
148 #[must_use]
150 pub const fn client_secret(&self) -> &ClientSecret {
151 &self.client_secret
152 }
153
154 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)]
194pub struct PkceFlow {
196 pub access_token: SecretAccessToken,
198 pub refresh_token: Option<RefreshToken>,
200}
201
202#[derive(Debug, thiserror::Error)]
204pub enum PkceFlowError {
205 #[error(transparent)]
207 PkceLogin(#[from] PkceLoginError),
208 #[error(transparent)]
210 Discovery(#[from] DiscoveryError),
211 #[error(transparent)]
213 Request(#[from] qcs_dependencies_client::reqwest::Error),
214}
215
216impl PkceFlow {
217 pub async fn new_login_flow(
223 cancel_token: CancellationToken,
224 auth_server: &AuthServer,
225 ) -> Result<Self, PkceFlowError> {
226 let issuer = auth_server.issuer.clone();
227
228 let client = default_http_client()?;
229 let discovery = oidc::fetch_discovery(&client, &issuer).await?;
230
231 let response = pkce_login(
232 cancel_token,
233 PkceLoginRequest {
234 client_id: auth_server.client_id.clone(),
235 redirect_port: None,
236 discovery,
237 scopes: auth_server.scopes.clone(),
238 },
239 )
240 .await?;
241
242 Ok(Self {
243 access_token: SecretAccessToken::from(response.access_token().secret().clone()),
244 refresh_token: response
245 .refresh_token()
246 .map(|rt| RefreshToken::new(SecretRefreshToken::from(rt.secret().clone()))),
247 })
248 }
249
250 pub async fn request_access_token(
256 &mut self,
257 auth_server: &AuthServer,
258 ) -> Result<SecretAccessToken, TokenError> {
259 if insecure_validate_token_exp(&self.access_token).is_ok() {
260 return Ok(self.access_token.clone());
261 }
262
263 if let Some(refresh_token) = &mut self.refresh_token {
264 let access_token = refresh_token.request_access_token(auth_server).await?;
265 self.access_token.clone_from(&access_token);
266 return Ok(access_token);
267 }
268
269 Err(TokenError::NoRefreshToken)
270 }
271}
272
273impl From<PkceFlow> for Credential {
274 fn from(value: PkceFlow) -> Self {
275 let mut token_payload = TokenPayload::default();
276 token_payload.access_token = Some(value.access_token);
277 token_payload.refresh_token = value.refresh_token.map(|rt| rt.refresh_token);
278
279 Self::TokenPayload(token_payload)
280 }
281}
282
283#[derive(Clone)]
284#[cfg_attr(feature = "python", derive(pyo3::FromPyObject, pyo3::IntoPyObject))]
285pub enum OAuthGrant {
288 RefreshToken(RefreshToken),
290 ClientCredentials(ClientCredentials),
292 ExternallyManaged(ExternallyManaged),
294 PkceFlow(PkceFlow),
296}
297
298impl From<ExternallyManaged> for OAuthGrant {
299 fn from(v: ExternallyManaged) -> Self {
300 Self::ExternallyManaged(v)
301 }
302}
303
304impl From<ClientCredentials> for OAuthGrant {
305 fn from(v: ClientCredentials) -> Self {
306 Self::ClientCredentials(v)
307 }
308}
309
310impl From<RefreshToken> for OAuthGrant {
311 fn from(v: RefreshToken) -> Self {
312 Self::RefreshToken(v)
313 }
314}
315
316impl From<PkceFlow> for OAuthGrant {
317 fn from(v: PkceFlow) -> Self {
318 Self::PkceFlow(v)
319 }
320}
321
322impl OAuthGrant {
323 async fn request_access_token(
325 &mut self,
326 auth_server: &AuthServer,
327 ) -> Result<SecretAccessToken, TokenError> {
328 match self {
329 Self::RefreshToken(tokens) => tokens.request_access_token(auth_server).await,
330 Self::ClientCredentials(tokens) => tokens.request_access_token(auth_server).await,
331 Self::ExternallyManaged(tokens) => tokens
332 .request_access_token(auth_server)
333 .await
334 .map_err(|e| TokenError::ExternallyManaged(e.to_string())),
335 Self::PkceFlow(tokens) => tokens.request_access_token(auth_server).await,
336 }
337 }
338}
339
340impl std::fmt::Debug for OAuthGrant {
341 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
342 match self {
343 Self::RefreshToken(_) => f.write_str("RefreshToken"),
344 Self::ClientCredentials(_) => f.write_str("ClientCredentials"),
345 Self::ExternallyManaged(_) => f.write_str("ExternallyManaged"),
346 Self::PkceFlow(_) => f.write_str("PkceTokens"),
347 }
348 }
349}
350
351#[derive(Clone)]
363#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
364#[cfg_attr(
365 feature = "python",
366 pyo3::pyclass(
367 module = "qcs_api_client_common._qcs_api_client_common.configuration",
368 frozen,
369 get_all,
370 from_py_object
371 )
372)]
373pub struct OAuthSession {
374 payload: OAuthGrant,
376 access_token: Option<SecretAccessToken>,
378 auth_server: AuthServer,
380}
381
382impl OAuthSession {
383 #[must_use]
388 pub const fn new(
389 payload: OAuthGrant,
390 auth_server: AuthServer,
391 access_token: Option<SecretAccessToken>,
392 ) -> Self {
393 Self {
394 payload,
395 access_token,
396 auth_server,
397 }
398 }
399
400 #[must_use]
405 pub const fn from_externally_managed(
406 tokens: ExternallyManaged,
407 auth_server: AuthServer,
408 access_token: Option<SecretAccessToken>,
409 ) -> Self {
410 Self::new(
411 OAuthGrant::ExternallyManaged(tokens),
412 auth_server,
413 access_token,
414 )
415 }
416
417 #[must_use]
422 pub const fn from_refresh_token(
423 tokens: RefreshToken,
424 auth_server: AuthServer,
425 access_token: Option<SecretAccessToken>,
426 ) -> Self {
427 Self::new(OAuthGrant::RefreshToken(tokens), auth_server, access_token)
428 }
429
430 #[must_use]
435 pub const fn from_client_credentials(
436 tokens: ClientCredentials,
437 auth_server: AuthServer,
438 access_token: Option<SecretAccessToken>,
439 ) -> Self {
440 Self::new(
441 OAuthGrant::ClientCredentials(tokens),
442 auth_server,
443 access_token,
444 )
445 }
446
447 #[must_use]
452 pub const fn from_pkce_flow(
453 flow: PkceFlow,
454 auth_server: AuthServer,
455 access_token: Option<SecretAccessToken>,
456 ) -> Self {
457 Self::new(OAuthGrant::PkceFlow(flow), auth_server, access_token)
458 }
459
460 pub fn access_token(&self) -> Result<&SecretAccessToken, TokenError> {
469 self.access_token.as_ref().ok_or(TokenError::NoAccessToken)
470 }
471
472 #[must_use]
474 pub const fn payload(&self) -> &OAuthGrant {
475 &self.payload
476 }
477
478 #[allow(clippy::missing_panics_doc)]
484 pub async fn request_access_token(&mut self) -> Result<&SecretAccessToken, TokenError> {
485 let access_token = self.payload.request_access_token(&self.auth_server).await?;
486 Ok(self.access_token.insert(access_token))
487 }
488
489 #[must_use]
491 pub const fn auth_server(&self) -> &AuthServer {
492 &self.auth_server
493 }
494
495 pub fn validate(&self) -> Result<SecretAccessToken, TokenError> {
503 let access_token = self.access_token()?;
504 insecure_validate_token_exp(access_token)?;
505 Ok(access_token.clone())
506 }
507}
508
509pub(crate) fn insecure_validate_token_exp(
513 access_token: &SecretAccessToken,
514) -> Result<(), TokenError> {
515 let placeholder_key = DecodingKey::from_secret(&[]);
516 let mut validation = Validation::new(Algorithm::RS256);
517 validation.validate_exp = true;
518 validation.leeway = 60;
519 validation.validate_aud = false;
520 validation.insecure_disable_signature_validation();
521
522 jsonwebtoken::decode::<toml::Value>(access_token.secret(), &placeholder_key, &validation)
523 .map(|_| ())
524 .map_err(TokenError::InvalidAccessToken)
525}
526
527impl std::fmt::Debug for OAuthSession {
528 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
529 let token_populated = if self.access_token.is_some() {
530 Some(())
531 } else {
532 None
533 };
534 f.debug_struct("OAuthSession")
535 .field("payload", &self.payload)
536 .field("access_token", &token_populated)
537 .field("auth_server", &self.auth_server)
538 .finish()
539 }
540}
541
542pub(crate) async fn persist_oauth_session(
557 oauth_session: &OAuthSession,
558 source: &ConfigSource,
559 credentials_name: &str,
560) -> Result<(), WriteError> {
561 let ConfigSource::File {
562 settings_path: _,
563 secrets_path,
564 } = source
565 else {
566 return Ok(());
567 };
568
569 let refresh_token = match &oauth_session.payload {
573 OAuthGrant::PkceFlow(payload) => payload.refresh_token.as_ref().map(|rt| &rt.refresh_token),
574 OAuthGrant::RefreshToken(payload) => Some(&payload.refresh_token),
575 OAuthGrant::ExternallyManaged(_) | OAuthGrant::ClientCredentials(_) => return Ok(()),
576 };
577
578 if Secrets::is_read_only(secrets_path).await? {
579 #[cfg(feature = "tracing")]
580 tracing::debug!(
581 "Skipping write of refreshed tokens to read-only secrets file: {:?}",
582 secrets_path
583 );
584 return Ok(());
585 }
586
587 let Ok(access_token) = oauth_session.access_token() else {
590 return Ok(());
591 };
592
593 let now = OffsetDateTime::now_utc();
594 Secrets::write_tokens(
595 secrets_path,
596 credentials_name,
597 refresh_token,
598 access_token,
599 now,
600 )
601 .await
602}
603
604#[derive(Clone, Debug)]
606#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
607#[cfg_attr(
608 feature = "python",
609 pyo3::pyclass(
610 module = "qcs_api_client_common._qcs_api_client_common.configuration",
611 frozen,
612 from_py_object
613 )
614)]
615pub struct TokenDispatcher {
616 lock: Arc<RwLock<OAuthSession>>,
617 refreshing: Arc<Mutex<bool>>,
618 notify_refreshed: Arc<Notify>,
619}
620
621impl From<OAuthSession> for TokenDispatcher {
622 fn from(value: OAuthSession) -> Self {
623 Self {
624 lock: Arc::new(RwLock::new(value)),
625 refreshing: Arc::new(Mutex::new(false)),
626 notify_refreshed: Arc::new(Notify::new()),
627 }
628 }
629}
630
631impl TokenDispatcher {
632 pub async fn use_tokens<F, O>(&self, f: F) -> O
642 where
643 F: FnOnce(&OAuthSession) -> O + Send,
644 {
645 let tokens = self.lock.read().await;
646 f(&tokens)
647 }
648
649 #[must_use]
651 pub async fn tokens(&self) -> OAuthSession {
652 self.use_tokens(Clone::clone).await
653 }
654
655 pub async fn refresh(
661 &self,
662 source: &ConfigSource,
663 credentials_name: &str,
664 ) -> Result<OAuthSession, TokenError> {
665 self.managed_refresh(Self::perform_refresh, source, credentials_name)
666 .await
667 }
668
669 pub async fn validate(&self) -> Result<SecretAccessToken, TokenError> {
677 self.use_tokens(OAuthSession::validate).await
678 }
679
680 async fn managed_refresh<F, Fut>(
683 &self,
684 refresh_fn: F,
685 source: &ConfigSource,
686 credentials_name: &str,
687 ) -> Result<OAuthSession, TokenError>
688 where
689 F: FnOnce(Arc<RwLock<OAuthSession>>) -> Fut + Send,
690 Fut: Future<Output = Result<OAuthSession, TokenError>> + Send,
691 {
692 let mut is_refreshing = self.refreshing.lock().await;
693
694 if *is_refreshing {
695 drop(is_refreshing);
696 self.notify_refreshed.notified().await;
697 return Ok(self.tokens().await);
698 }
699
700 *is_refreshing = true;
701 drop(is_refreshing);
702
703 let oauth_session = refresh_fn(self.lock.clone()).await?;
704
705 let write_result = persist_oauth_session(&oauth_session, source, credentials_name).await;
706
707 *self.refreshing.lock().await = false;
709 self.notify_refreshed.notify_waiters();
710
711 if let Err(error) = write_result {
713 return Err(TokenError::Write {
714 error,
715 oauth_session: Box::new(oauth_session),
716 });
717 }
718
719 Ok(oauth_session)
720 }
721
722 async fn perform_refresh(lock: Arc<RwLock<OAuthSession>>) -> Result<OAuthSession, TokenError> {
729 let mut credentials = lock.write().await;
730 credentials.request_access_token().await?;
731 Ok(credentials.clone())
732 }
733}
734
735pub(crate) type RefreshResult =
736 Pin<Box<dyn Future<Output = Result<String, Box<dyn std::error::Error + Send + Sync>>> + Send>>;
737
738pub type RefreshFunction = Box<dyn (Fn(AuthServer) -> RefreshResult) + Send + Sync>;
740
741#[derive(Clone)]
746#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
747#[cfg_attr(
748 feature = "python",
749 pyo3::pyclass(
750 module = "qcs_api_client_common._qcs_api_client_common.configuration",
751 frozen,
752 from_py_object
753 )
754)]
755pub struct ExternallyManaged {
756 refresh_function: Arc<RefreshFunction>,
757}
758
759impl ExternallyManaged {
760 pub fn new(
785 refresh_function: impl Fn(AuthServer) -> RefreshResult + Send + Sync + 'static,
786 ) -> Self {
787 Self {
788 refresh_function: Arc::new(Box::new(refresh_function)),
789 }
790 }
791
792 pub fn from_async<F, Fut>(refresh_function: F) -> Self
825 where
826 F: Fn(AuthServer) -> Fut + Send + Sync + 'static,
827 Fut: Future<Output = Result<String, Box<dyn std::error::Error + Send + Sync>>>
828 + Send
829 + 'static,
830 {
831 Self {
832 refresh_function: Arc::new(Box::new(move |auth_server| {
833 Box::pin(refresh_function(auth_server))
834 })),
835 }
836 }
837
838 pub fn from_sync(
869 refresh_function: impl Fn(
870 AuthServer,
871 ) -> Result<String, Box<dyn std::error::Error + Send + Sync>>
872 + Send
873 + Sync
874 + 'static,
875 ) -> Self {
876 Self {
877 refresh_function: Arc::new(Box::new(move |auth_server| {
878 let result = refresh_function(auth_server);
879 Box::pin(async move { result })
880 })),
881 }
882 }
883
884 pub async fn request_access_token(
890 &self,
891 auth_server: &AuthServer,
892 ) -> Result<SecretAccessToken, Box<dyn std::error::Error + Send + Sync>> {
893 (self.refresh_function)(auth_server.clone())
894 .await
895 .map(SecretAccessToken::from)
896 }
897}
898
899impl std::fmt::Debug for ExternallyManaged {
900 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
901 f.debug_struct("ExternallyManaged")
902 .field(
903 "refresh_function",
904 &"Fn() -> Pin<Box<dyn Future<Output = Result<String, TokenError>> + Send>>",
905 )
906 .finish()
907 }
908}
909
910#[derive(Debug, Serialize, Deserialize)]
911pub(super) struct TokenRefreshRequest<'a> {
912 grant_type: &'static str,
913 client_id: &'a str,
914 refresh_token: &'a str,
915}
916
917impl<'a> TokenRefreshRequest<'a> {
918 pub(super) const fn new(client_id: &'a str, refresh_token: &'a str) -> Self {
919 Self {
920 grant_type: "refresh_token",
921 client_id,
922 refresh_token,
923 }
924 }
925}
926
927#[derive(Debug, Serialize, Deserialize)]
928pub(super) struct ClientCredentialsRequest {
929 grant_type: &'static str,
930 scope: Option<&'static str>,
931}
932
933impl ClientCredentialsRequest {
934 pub(super) const fn new(scope: Option<&'static str>) -> Self {
935 Self {
936 grant_type: "client_credentials",
937 scope,
938 }
939 }
940}
941
942#[derive(Deserialize, Debug, Serialize)]
943pub(super) struct RefreshTokenResponse {
944 pub(super) refresh_token: Option<SecretRefreshToken>,
945 pub(super) access_token: SecretAccessToken,
946}
947
948#[async_trait::async_trait]
950pub trait TokenRefresher: Clone + std::fmt::Debug + Send {
951 type Error;
954
955 async fn validated_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
957
958 async fn get_access_token(&self) -> Result<Option<SecretAccessToken>, Self::Error>;
960
961 async fn refresh_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
963
964 #[cfg(feature = "tracing")]
966 fn base_url(&self) -> &str;
967
968 #[cfg(feature = "tracing-config")]
970 fn tracing_configuration(&self) -> Option<&TracingConfiguration>;
971
972 #[cfg(feature = "tracing")]
975 #[allow(clippy::needless_return)]
976 fn should_trace(&self, url: &UrlPatternMatchInput) -> bool {
977 #[cfg(not(feature = "tracing-config"))]
978 {
979 let _ = url;
980 return true;
981 }
982
983 #[cfg(feature = "tracing-config")]
984 self.tracing_configuration()
985 .is_none_or(|config| config.is_enabled(url))
986 }
987}
988
989#[async_trait::async_trait]
990impl TokenRefresher for ClientConfiguration {
991 type Error = TokenError;
992
993 async fn validated_access_token(&self) -> Result<SecretAccessToken, Self::Error> {
994 self.get_bearer_access_token().await
995 }
996
997 async fn refresh_access_token(&self) -> Result<SecretAccessToken, Self::Error> {
998 match self.refresh().await {
999 Ok(session) => Ok(session.access_token()?.clone()),
1000 Err(TokenError::Write {
1001 error: _error,
1002 oauth_session,
1003 }) => {
1004 #[cfg(feature = "tracing")]
1006 tracing::warn!(
1007 "Token refresh succeeded but failed to persist: {_error}. Returning access token from error.",
1008 );
1009 Ok(oauth_session.access_token()?.clone())
1010 }
1011 Err(e) => Err(e),
1012 }
1013 }
1014
1015 async fn get_access_token(&self) -> Result<Option<SecretAccessToken>, Self::Error> {
1016 Ok(Some(self.oauth_session().await?.access_token()?.clone()))
1017 }
1018
1019 #[cfg(feature = "tracing")]
1020 fn base_url(&self) -> &str {
1021 &self.grpc_api_url
1022 }
1023
1024 #[cfg(feature = "tracing-config")]
1025 fn tracing_configuration(&self) -> Option<&TracingConfiguration> {
1026 self.tracing_configuration.as_ref()
1027 }
1028}
1029
1030pub(super) fn default_http_client()
1032-> Result<qcs_dependencies_client::reqwest::Client, qcs_dependencies_client::reqwest::Error> {
1033 qcs_dependencies_client::reqwest::Client::builder()
1034 .timeout(std::time::Duration::from_secs(10))
1035 .build()
1036}
1037
1038#[cfg(test)]
1039mod test {
1040 #![allow(clippy::result_large_err, reason = "happens in figment tests")]
1041
1042 use std::time::Duration;
1043
1044 use super::*;
1045 use httpmock::prelude::*;
1046 use rstest::rstest;
1047 use time::format_description::well_known::Rfc3339;
1048 use tokio::time::Instant;
1049 use toml_edit::DocumentMut;
1050
1051 #[tokio::test]
1052 async fn test_tokens_blocked_during_refresh() {
1053 let mock_server = MockServer::start_async().await;
1054
1055 let oidc_mock = mock_server
1056 .mock_async(|when, then| {
1057 when.method(GET).path("/.well-known/openid-configuration");
1058 then.status(200)
1059 .json_body_obj(&oidc::Discovery::new_for_test(
1060 mock_server.base_url().parse().unwrap(),
1061 ));
1062 })
1063 .await;
1064
1065 let issuer_mock = mock_server
1066 .mock_async(|when, then| {
1067 when.method(POST).path("/v1/token");
1068
1069 then.status(200)
1070 .delay(Duration::from_secs(3))
1071 .json_body_obj(&RefreshTokenResponse {
1072 access_token: SecretAccessToken::from("new_access"),
1073 refresh_token: Some(SecretRefreshToken::from("new_refresh")),
1074 });
1075 })
1076 .await;
1077
1078 let original_tokens = OAuthSession::from_refresh_token(
1079 RefreshToken::new(SecretRefreshToken::from("refresh")),
1080 AuthServer {
1081 client_id: "client_id".to_string(),
1082 issuer: mock_server.base_url(),
1083 scopes: None,
1084 },
1085 None,
1086 );
1087 let dispatcher: TokenDispatcher = original_tokens.clone().into();
1088 let dispatcher_clone1 = dispatcher.clone();
1089 let dispatcher_clone2 = dispatcher.clone();
1090
1091 let refresh_duration = Duration::from_secs(3);
1092
1093 let start_write = Instant::now();
1094 let write_future = tokio::spawn(async move {
1095 dispatcher_clone1
1096 .refresh(&ConfigSource::Default, "")
1097 .await
1098 .unwrap()
1099 });
1100
1101 let start_read = Instant::now();
1102 let read_future = tokio::spawn(async move { dispatcher_clone2.tokens().await });
1103
1104 let _ = write_future.await.unwrap();
1105 let read_result = read_future.await.unwrap();
1106
1107 let write_duration = start_write.elapsed();
1108 let read_duration = start_read.elapsed();
1109
1110 oidc_mock.assert_async().await;
1111 issuer_mock.assert_async().await;
1112
1113 assert!(
1114 write_duration >= refresh_duration,
1115 "Write operation did not take enough time"
1116 );
1117 assert!(
1118 read_duration >= refresh_duration,
1119 "Read operation was not blocked by the write operation"
1120 );
1121 assert_eq!(
1122 read_result.access_token.unwrap(),
1123 SecretAccessToken::from("new_access")
1124 );
1125 if let OAuthGrant::RefreshToken(payload) = read_result.payload {
1126 assert_eq!(
1127 payload.refresh_token,
1128 SecretRefreshToken::from("new_refresh")
1129 );
1130 } else {
1131 panic!(
1132 "Expected RefreshToken payload, got {:?}",
1133 read_result.payload
1134 );
1135 }
1136 }
1137
1138 #[tokio::test]
1142 async fn test_refresh_token_request_rejected_by_auth_server() {
1143 let mock_server = MockServer::start_async().await;
1144
1145 let oidc_mock = mock_server
1146 .mock_async(|when, then| {
1147 when.method(GET).path("/.well-known/openid-configuration");
1148 then.status(200)
1149 .json_body_obj(&oidc::Discovery::new_for_test(
1150 mock_server.base_url().parse().unwrap(),
1151 ));
1152 })
1153 .await;
1154
1155 let issuer_mock = mock_server
1156 .mock_async(|when, then| {
1157 when.method(POST).path("/v1/token");
1158 then.status(400).json_body_obj(&serde_json::json!({
1159 "error": "invalid_grant",
1160 "error_description": "Unknown or invalid refresh token.",
1161 }));
1162 })
1163 .await;
1164
1165 let mut refresh_token = RefreshToken::new(SecretRefreshToken::from("revoked_refresh"));
1166 let auth_server = AuthServer {
1167 client_id: "client_id".to_string(),
1168 issuer: mock_server.base_url(),
1169 scopes: None,
1170 };
1171
1172 let result = refresh_token.request_access_token(&auth_server).await;
1173
1174 oidc_mock.assert_async().await;
1175 issuer_mock.assert_async().await;
1176
1177 assert!(
1178 result.is_err(),
1179 "a rejected refresh token request should be an error, got {result:?}"
1180 );
1181 }
1182
1183 #[rstest]
1184 fn test_qcs_secrets_readonly(
1185 #[values(
1186 (Some("TRUE"), true),
1187 (Some("tRue"), true),
1188 (Some("true"), true),
1189 (Some("YES"), true),
1190 (Some("yEs"), true),
1191 (Some("yes"), true),
1192 (Some("1"), true),
1193 (Some("2"), false),
1194 (Some("other"), false),
1195 (Some(""), false),
1196 (None, false),
1197 )]
1198 read_only_values: (Option<&str>, bool),
1199 #[values(true, false)] read_only_perm: bool,
1200 ) {
1201 let (maybe_read_only_env, env_is_read_only) = read_only_values;
1202 let expected_update = !env_is_read_only && !read_only_perm;
1203 figment::Jail::expect_with(|jail| {
1204 let profile_name = "test";
1205 let initial_access_token = "initial_access_token";
1206 let initial_refresh_token = "initial_refresh_token";
1207
1208 let initial_secrets_file_contents = format!(
1209 r#"
1210[credentials]
1211[credentials.{profile_name}]
1212[credentials.{profile_name}.token_payload]
1213access_token = "{initial_access_token}"
1214expires_in = 3600
1215id_token = "id_token"
1216refresh_token = "{initial_refresh_token}"
1217scope = "offline_access openid profile email"
1218token_type = "Bearer"
1219updated_at = "2024-01-01T00:00:00Z"
1220"#
1221 );
1222
1223 jail.clear_env();
1225
1226 let secrets_path = "secrets.toml";
1228 jail.create_file(secrets_path, initial_secrets_file_contents.as_str())
1229 .expect("should create test secrets.toml");
1230
1231 if read_only_perm {
1232 let mut permissions = std::fs::metadata(secrets_path)
1233 .expect("Should be able to get file metadata")
1234 .permissions();
1235 permissions.set_readonly(true);
1236 std::fs::set_permissions(secrets_path, permissions)
1237 .expect("Should be able to set file permissions");
1238 }
1239
1240 let rt = tokio::runtime::Runtime::new().unwrap();
1241 rt.block_on(async {
1242 let mock_server = MockServer::start_async().await;
1243
1244 let oidc_mock = mock_server
1245 .mock_async(|when, then| {
1246 when.method(GET).path("/.well-known/openid-configuration");
1247 then.status(200)
1248 .json_body_obj(&oidc::Discovery::new_for_test(mock_server.base_url().parse().unwrap()));
1249 })
1250 .await;
1251
1252 let new_access_token = SecretAccessToken::from("new_access_token");
1254 let issuer_mock = mock_server
1255 .mock_async(|when, then| {
1256 when.method(POST).path("/v1/token");
1257 then.status(200).json_body_obj(&RefreshTokenResponse {
1258 access_token: new_access_token.clone(),
1259 refresh_token: Some(SecretRefreshToken::from(initial_refresh_token)),
1260 });
1261 })
1262 .await;
1263
1264 let original_tokens = OAuthSession::from_refresh_token(
1266 RefreshToken::new(SecretRefreshToken::from(initial_refresh_token)),
1267 AuthServer { client_id: "client_id".to_string(), issuer: mock_server.base_url(), scopes: None },
1268 Some(SecretAccessToken::from(initial_refresh_token)),
1269 );
1270 let dispatcher: TokenDispatcher = original_tokens.into();
1271
1272 jail.set_env("QCS_SECRETS_FILE_PATH", "secrets.toml");
1274 jail.set_env("QCS_PROFILE_NAME", "test");
1275 if let Some(read_only_env) = maybe_read_only_env {
1276 jail.set_env("QCS_SECRETS_READ_ONLY", read_only_env);
1277 }
1278
1279 let before_refresh = OffsetDateTime::now_utc();
1280
1281 dispatcher
1282 .refresh(
1283 &ConfigSource::File {
1284 settings_path: "".into(),
1285 secrets_path: "secrets.toml".into(),
1286 },
1287 profile_name,
1288 )
1289 .await
1290 .unwrap();
1291
1292 oidc_mock.assert_async().await;
1293 issuer_mock.assert_async().await;
1294
1295 let content = std::fs::read_to_string("secrets.toml").unwrap();
1297 if !expected_update {
1298 assert!(
1299 content.eq(initial_secrets_file_contents.as_str()),
1300 "File should not be updated when QCS_SECRETS_READ_ONLY is set or file permissions are read-only"
1301 );
1302 return;
1303 }
1304
1305 let mut toml = std::fs::read_to_string(secrets_path)
1307 .unwrap()
1308 .parse::<DocumentMut>()
1309 .unwrap();
1310
1311 let token_payload = toml
1312 .get_mut("credentials")
1313 .and_then(|credentials| {
1314 credentials.get_mut(profile_name)?.get_mut("token_payload")
1315 })
1316 .expect("Should be able to get token_payload table");
1317
1318 let access_token = token_payload.get("access_token").unwrap().as_str().map(str::to_string).map(SecretAccessToken::from);
1319
1320 assert_eq!(
1321 access_token,
1322 Some(new_access_token)
1323 );
1324
1325 assert!(
1326 OffsetDateTime::parse(
1327 token_payload.get("updated_at").unwrap().as_str().unwrap(),
1328 &Rfc3339
1329 )
1330 .unwrap()
1331 > before_refresh
1332 );
1333
1334 let content = std::fs::read_to_string("secrets.toml").unwrap();
1335 assert!(
1336 content.contains("new_access_token"),
1337 "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"
1338 );
1339 });
1340 Ok(())
1341 });
1342 }
1343
1344 #[test]
1347 fn test_refresh_token_grant_persists_rotated_refresh_token() {
1348 let initial_refresh_token = "initial_refresh_token";
1349 let rotated_refresh_token = "rotated_refresh_token";
1350 let new_access_token = "new_access_token";
1351
1352 figment::Jail::expect_with(|jail| {
1353 jail.clear_env();
1354
1355 let secrets_path = "secrets.toml";
1356 let initial_secrets_file_contents = format!(
1357 r#"
1358[credentials]
1359[credentials.test]
1360[credentials.test.token_payload]
1361access_token = "initial_access_token"
1362refresh_token = "{initial_refresh_token}"
1363updated_at = "2024-01-01T00:00:00Z"
1364"#
1365 );
1366 jail.create_file(secrets_path, &initial_secrets_file_contents)
1367 .expect("should create test secrets.toml");
1368
1369 let rt = tokio::runtime::Runtime::new().unwrap();
1370 rt.block_on(async {
1371 let mock_server = MockServer::start_async().await;
1372 let oidc_mock = mock_server
1373 .mock_async(|when, then| {
1374 when.method(GET).path("/.well-known/openid-configuration");
1375 then.status(200)
1376 .json_body_obj(&oidc::Discovery::new_for_test(
1377 mock_server.base_url().parse().unwrap(),
1378 ));
1379 })
1380 .await;
1381 let issuer_mock = mock_server
1382 .mock_async(|when, then| {
1383 when.method(POST).path("/v1/token");
1384 then.status(200).json_body_obj(&RefreshTokenResponse {
1385 access_token: SecretAccessToken::from(new_access_token),
1386 refresh_token: Some(SecretRefreshToken::from(rotated_refresh_token)),
1387 });
1388 })
1389 .await;
1390
1391 let dispatcher: TokenDispatcher = OAuthSession::from_refresh_token(
1392 RefreshToken::new(SecretRefreshToken::from(initial_refresh_token)),
1393 AuthServer {
1394 client_id: "client_id".to_string(),
1395 issuer: mock_server.base_url(),
1396 scopes: None,
1397 },
1398 Some(SecretAccessToken::from("initial_access_token")),
1399 )
1400 .into();
1401
1402 dispatcher
1403 .refresh(
1404 &ConfigSource::File {
1405 settings_path: "".into(),
1406 secrets_path: secrets_path.into(),
1407 },
1408 "test",
1409 )
1410 .await
1411 .expect("refresh should succeed");
1412
1413 oidc_mock.assert_async().await;
1414 issuer_mock.assert_async().await;
1415 });
1416
1417 let Credential::TokenPayload(payload) = Secrets::load_from_path(&secrets_path.into())
1419 .expect("should load secrets")
1420 .credentials
1421 .remove("test")
1422 .expect("should have test credentials")
1423 else {
1424 panic!("expected a token payload credential");
1425 };
1426 assert_eq!(
1427 payload.refresh_token.unwrap(),
1428 SecretRefreshToken::from(rotated_refresh_token),
1429 "rotated refresh token should be persisted to the secrets file"
1430 );
1431 assert_eq!(
1432 payload.access_token.unwrap(),
1433 SecretAccessToken::from(new_access_token),
1434 "new access token should be persisted to the secrets file"
1435 );
1436
1437 Ok(())
1438 });
1439 }
1440
1441 #[test]
1442 fn test_auth_session_debug_fmt() {
1443 let session = OAuthSession {
1444 payload: OAuthGrant::ClientCredentials(ClientCredentials::new(
1445 "hidden_id",
1446 "hidden_secret",
1447 )),
1448 access_token: Some(SecretAccessToken::from("token")),
1449 auth_server: AuthServer {
1450 client_id: "some_id".into(),
1451 issuer: "some_url".into(),
1452 scopes: None,
1453 },
1454 };
1455
1456 assert_eq!(
1457 "OAuthSession { payload: ClientCredentials, access_token: Some(()), auth_server: AuthServer { client_id: \"some_id\", issuer: \"some_url\", scopes: None } }",
1458 &format!("{session:?}")
1459 );
1460 }
1461}