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, RedirectBinding, pkce_login},
21 secrets::{Credential, SecretAccessToken, SecretRefreshToken, TokenPayload},
22};
23#[cfg(feature = "tracing-config")]
24use crate::tracing_configuration::TracingConfiguration;
25#[cfg(feature = "tracing")]
26use urlpattern::UrlPatternMatchInput;
27
28pub use super::secret_string::ClientSecret;
29
30#[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 Self::new_login_flow_with_redirect(cancel_token, auth_server, RedirectBinding::default())
227 .await
228 }
229
230 pub(crate) async fn new_login_flow_with_redirect(
232 cancel_token: CancellationToken,
233 auth_server: &AuthServer,
234 redirect: RedirectBinding,
235 ) -> Result<Self, PkceFlowError> {
236 let issuer = auth_server.issuer.clone();
237
238 let client = default_http_client()?;
239 let discovery = oidc::fetch_discovery(&client, &issuer).await?;
240
241 let response = pkce_login(
242 cancel_token,
243 PkceLoginRequest {
244 client_id: auth_server.client_id.clone(),
245 redirect,
246 discovery,
247 scopes: auth_server.scopes.clone(),
248 },
249 )
250 .await?;
251
252 Ok(Self {
253 access_token: SecretAccessToken::from(response.access_token().secret().clone()),
254 refresh_token: response
255 .refresh_token()
256 .map(|rt| RefreshToken::new(SecretRefreshToken::from(rt.secret().clone()))),
257 })
258 }
259
260 pub async fn request_access_token(
266 &mut self,
267 auth_server: &AuthServer,
268 ) -> Result<SecretAccessToken, TokenError> {
269 if insecure_validate_token_exp(&self.access_token).is_ok() {
270 return Ok(self.access_token.clone());
271 }
272
273 if let Some(refresh_token) = &mut self.refresh_token {
274 let access_token = refresh_token.request_access_token(auth_server).await?;
275 self.access_token.clone_from(&access_token);
276 return Ok(access_token);
277 }
278
279 Err(TokenError::NoRefreshToken)
280 }
281}
282
283impl From<PkceFlow> for Credential {
284 fn from(value: PkceFlow) -> Self {
285 let mut token_payload = TokenPayload::default();
286 token_payload.access_token = Some(value.access_token);
287 token_payload.refresh_token = value.refresh_token.map(|rt| rt.refresh_token);
288
289 Self::TokenPayload(token_payload)
290 }
291}
292
293#[derive(Clone)]
294#[cfg_attr(feature = "python", derive(pyo3::FromPyObject, pyo3::IntoPyObject))]
295pub enum OAuthGrant {
298 RefreshToken(RefreshToken),
300 ClientCredentials(ClientCredentials),
302 ExternallyManaged(ExternallyManaged),
304 PkceFlow(PkceFlow),
306}
307
308impl From<ExternallyManaged> for OAuthGrant {
309 fn from(v: ExternallyManaged) -> Self {
310 Self::ExternallyManaged(v)
311 }
312}
313
314impl From<ClientCredentials> for OAuthGrant {
315 fn from(v: ClientCredentials) -> Self {
316 Self::ClientCredentials(v)
317 }
318}
319
320impl From<RefreshToken> for OAuthGrant {
321 fn from(v: RefreshToken) -> Self {
322 Self::RefreshToken(v)
323 }
324}
325
326impl From<PkceFlow> for OAuthGrant {
327 fn from(v: PkceFlow) -> Self {
328 Self::PkceFlow(v)
329 }
330}
331
332impl OAuthGrant {
333 async fn request_access_token(
335 &mut self,
336 auth_server: &AuthServer,
337 ) -> Result<SecretAccessToken, TokenError> {
338 match self {
339 Self::RefreshToken(tokens) => tokens.request_access_token(auth_server).await,
340 Self::ClientCredentials(tokens) => tokens.request_access_token(auth_server).await,
341 Self::ExternallyManaged(tokens) => tokens
342 .request_access_token(auth_server)
343 .await
344 .map_err(|e| TokenError::ExternallyManaged(e.to_string())),
345 Self::PkceFlow(tokens) => tokens.request_access_token(auth_server).await,
346 }
347 }
348}
349
350impl std::fmt::Debug for OAuthGrant {
351 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
352 match self {
353 Self::RefreshToken(_) => f.write_str("RefreshToken"),
354 Self::ClientCredentials(_) => f.write_str("ClientCredentials"),
355 Self::ExternallyManaged(_) => f.write_str("ExternallyManaged"),
356 Self::PkceFlow(_) => f.write_str("PkceTokens"),
357 }
358 }
359}
360
361#[derive(Clone)]
373#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
374#[cfg_attr(
375 feature = "python",
376 pyo3::pyclass(
377 module = "qcs_api_client_common._qcs_api_client_common.configuration",
378 frozen,
379 get_all,
380 from_py_object
381 )
382)]
383pub struct OAuthSession {
384 payload: OAuthGrant,
386 access_token: Option<SecretAccessToken>,
388 auth_server: AuthServer,
390}
391
392impl OAuthSession {
393 #[must_use]
398 pub const fn new(
399 payload: OAuthGrant,
400 auth_server: AuthServer,
401 access_token: Option<SecretAccessToken>,
402 ) -> Self {
403 Self {
404 payload,
405 access_token,
406 auth_server,
407 }
408 }
409
410 #[must_use]
415 pub const fn from_externally_managed(
416 tokens: ExternallyManaged,
417 auth_server: AuthServer,
418 access_token: Option<SecretAccessToken>,
419 ) -> Self {
420 Self::new(
421 OAuthGrant::ExternallyManaged(tokens),
422 auth_server,
423 access_token,
424 )
425 }
426
427 #[must_use]
432 pub const fn from_refresh_token(
433 tokens: RefreshToken,
434 auth_server: AuthServer,
435 access_token: Option<SecretAccessToken>,
436 ) -> Self {
437 Self::new(OAuthGrant::RefreshToken(tokens), auth_server, access_token)
438 }
439
440 #[must_use]
445 pub const fn from_client_credentials(
446 tokens: ClientCredentials,
447 auth_server: AuthServer,
448 access_token: Option<SecretAccessToken>,
449 ) -> Self {
450 Self::new(
451 OAuthGrant::ClientCredentials(tokens),
452 auth_server,
453 access_token,
454 )
455 }
456
457 #[must_use]
462 pub const fn from_pkce_flow(
463 flow: PkceFlow,
464 auth_server: AuthServer,
465 access_token: Option<SecretAccessToken>,
466 ) -> Self {
467 Self::new(OAuthGrant::PkceFlow(flow), auth_server, access_token)
468 }
469
470 pub fn access_token(&self) -> Result<&SecretAccessToken, TokenError> {
479 self.access_token.as_ref().ok_or(TokenError::NoAccessToken)
480 }
481
482 #[must_use]
484 pub const fn payload(&self) -> &OAuthGrant {
485 &self.payload
486 }
487
488 #[allow(clippy::missing_panics_doc)]
494 pub async fn request_access_token(&mut self) -> Result<&SecretAccessToken, TokenError> {
495 let access_token = self.payload.request_access_token(&self.auth_server).await?;
496 Ok(self.access_token.insert(access_token))
497 }
498
499 #[must_use]
501 pub const fn auth_server(&self) -> &AuthServer {
502 &self.auth_server
503 }
504
505 pub fn validate(&self) -> Result<SecretAccessToken, TokenError> {
513 let access_token = self.access_token()?;
514 insecure_validate_token_exp(access_token)?;
515 Ok(access_token.clone())
516 }
517}
518
519pub(crate) fn insecure_validate_token_exp(
523 access_token: &SecretAccessToken,
524) -> Result<(), TokenError> {
525 let placeholder_key = DecodingKey::from_secret(&[]);
526 let mut validation = Validation::new(Algorithm::RS256);
527 validation.validate_exp = true;
528 validation.leeway = 60;
529 validation.validate_aud = false;
530 validation.insecure_disable_signature_validation();
531
532 jsonwebtoken::decode::<toml::Value>(access_token.secret(), &placeholder_key, &validation)
533 .map(|_| ())
534 .map_err(TokenError::InvalidAccessToken)
535}
536
537impl std::fmt::Debug for OAuthSession {
538 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
539 let token_populated = if self.access_token.is_some() {
540 Some(())
541 } else {
542 None
543 };
544 f.debug_struct("OAuthSession")
545 .field("payload", &self.payload)
546 .field("access_token", &token_populated)
547 .field("auth_server", &self.auth_server)
548 .finish()
549 }
550}
551
552pub(crate) async fn persist_oauth_session(
567 oauth_session: &OAuthSession,
568 source: &ConfigSource,
569 credentials_name: &str,
570) -> Result<(), WriteError> {
571 let ConfigSource::File {
572 settings_path: _,
573 secrets_path,
574 } = source
575 else {
576 return Ok(());
577 };
578
579 let refresh_token = match &oauth_session.payload {
583 OAuthGrant::PkceFlow(payload) => payload.refresh_token.as_ref().map(|rt| &rt.refresh_token),
584 OAuthGrant::RefreshToken(payload) => Some(&payload.refresh_token),
585 OAuthGrant::ExternallyManaged(_) | OAuthGrant::ClientCredentials(_) => return Ok(()),
586 };
587
588 if Secrets::is_read_only(secrets_path).await? {
589 #[cfg(feature = "tracing")]
590 tracing::debug!(
591 "Skipping write of refreshed tokens to read-only secrets file: {:?}",
592 secrets_path
593 );
594 return Ok(());
595 }
596
597 let Ok(access_token) = oauth_session.access_token() else {
600 return Ok(());
601 };
602
603 let now = OffsetDateTime::now_utc();
604 Secrets::write_tokens(
605 secrets_path,
606 credentials_name,
607 refresh_token,
608 access_token,
609 now,
610 )
611 .await
612}
613
614#[derive(Clone, Debug)]
616#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
617#[cfg_attr(
618 feature = "python",
619 pyo3::pyclass(
620 module = "qcs_api_client_common._qcs_api_client_common.configuration",
621 frozen,
622 from_py_object
623 )
624)]
625pub struct TokenDispatcher {
626 lock: Arc<RwLock<OAuthSession>>,
627 refreshing: Arc<Mutex<bool>>,
628 notify_refreshed: Arc<Notify>,
629}
630
631impl From<OAuthSession> for TokenDispatcher {
632 fn from(value: OAuthSession) -> Self {
633 Self {
634 lock: Arc::new(RwLock::new(value)),
635 refreshing: Arc::new(Mutex::new(false)),
636 notify_refreshed: Arc::new(Notify::new()),
637 }
638 }
639}
640
641impl TokenDispatcher {
642 pub async fn use_tokens<F, O>(&self, f: F) -> O
652 where
653 F: FnOnce(&OAuthSession) -> O + Send,
654 {
655 let tokens = self.lock.read().await;
656 f(&tokens)
657 }
658
659 #[must_use]
661 pub async fn tokens(&self) -> OAuthSession {
662 self.use_tokens(Clone::clone).await
663 }
664
665 pub async fn refresh(
671 &self,
672 source: &ConfigSource,
673 credentials_name: &str,
674 ) -> Result<OAuthSession, TokenError> {
675 self.managed_refresh(Self::perform_refresh, source, credentials_name)
676 .await
677 }
678
679 pub async fn validate(&self) -> Result<SecretAccessToken, TokenError> {
687 self.use_tokens(OAuthSession::validate).await
688 }
689
690 async fn managed_refresh<F, Fut>(
693 &self,
694 refresh_fn: F,
695 source: &ConfigSource,
696 credentials_name: &str,
697 ) -> Result<OAuthSession, TokenError>
698 where
699 F: FnOnce(Arc<RwLock<OAuthSession>>) -> Fut + Send,
700 Fut: Future<Output = Result<OAuthSession, TokenError>> + Send,
701 {
702 let mut is_refreshing = self.refreshing.lock().await;
703
704 if *is_refreshing {
705 drop(is_refreshing);
706 self.notify_refreshed.notified().await;
707 return Ok(self.tokens().await);
708 }
709
710 *is_refreshing = true;
711 drop(is_refreshing);
712
713 let oauth_session = refresh_fn(self.lock.clone()).await?;
714
715 let write_result = persist_oauth_session(&oauth_session, source, credentials_name).await;
716
717 *self.refreshing.lock().await = false;
719 self.notify_refreshed.notify_waiters();
720
721 if let Err(error) = write_result {
723 return Err(TokenError::Write {
724 error,
725 oauth_session: Box::new(oauth_session),
726 });
727 }
728
729 Ok(oauth_session)
730 }
731
732 async fn perform_refresh(lock: Arc<RwLock<OAuthSession>>) -> Result<OAuthSession, TokenError> {
739 let mut credentials = lock.write().await;
740 credentials.request_access_token().await?;
741 Ok(credentials.clone())
742 }
743}
744
745pub(crate) type RefreshResult =
746 Pin<Box<dyn Future<Output = Result<String, Box<dyn std::error::Error + Send + Sync>>> + Send>>;
747
748pub type RefreshFunction = Box<dyn (Fn(AuthServer) -> RefreshResult) + Send + Sync>;
750
751#[derive(Clone)]
756#[cfg_attr(feature = "stubs", gen_stub_pyclass)]
757#[cfg_attr(
758 feature = "python",
759 pyo3::pyclass(
760 module = "qcs_api_client_common._qcs_api_client_common.configuration",
761 frozen,
762 from_py_object
763 )
764)]
765pub struct ExternallyManaged {
766 refresh_function: Arc<RefreshFunction>,
767}
768
769impl ExternallyManaged {
770 pub fn new(
795 refresh_function: impl Fn(AuthServer) -> RefreshResult + Send + Sync + 'static,
796 ) -> Self {
797 Self {
798 refresh_function: Arc::new(Box::new(refresh_function)),
799 }
800 }
801
802 pub fn from_async<F, Fut>(refresh_function: F) -> Self
835 where
836 F: Fn(AuthServer) -> Fut + Send + Sync + 'static,
837 Fut: Future<Output = Result<String, Box<dyn std::error::Error + Send + Sync>>>
838 + Send
839 + 'static,
840 {
841 Self {
842 refresh_function: Arc::new(Box::new(move |auth_server| {
843 Box::pin(refresh_function(auth_server))
844 })),
845 }
846 }
847
848 pub fn from_sync(
879 refresh_function: impl Fn(
880 AuthServer,
881 ) -> Result<String, Box<dyn std::error::Error + Send + Sync>>
882 + Send
883 + Sync
884 + 'static,
885 ) -> Self {
886 Self {
887 refresh_function: Arc::new(Box::new(move |auth_server| {
888 let result = refresh_function(auth_server);
889 Box::pin(async move { result })
890 })),
891 }
892 }
893
894 pub async fn request_access_token(
900 &self,
901 auth_server: &AuthServer,
902 ) -> Result<SecretAccessToken, Box<dyn std::error::Error + Send + Sync>> {
903 (self.refresh_function)(auth_server.clone())
904 .await
905 .map(SecretAccessToken::from)
906 }
907}
908
909impl std::fmt::Debug for ExternallyManaged {
910 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
911 f.debug_struct("ExternallyManaged")
912 .field(
913 "refresh_function",
914 &"Fn() -> Pin<Box<dyn Future<Output = Result<String, TokenError>> + Send>>",
915 )
916 .finish()
917 }
918}
919
920#[derive(Debug, Serialize, Deserialize)]
921pub(super) struct TokenRefreshRequest<'a> {
922 grant_type: &'static str,
923 client_id: &'a str,
924 refresh_token: &'a str,
925}
926
927impl<'a> TokenRefreshRequest<'a> {
928 pub(super) const fn new(client_id: &'a str, refresh_token: &'a str) -> Self {
929 Self {
930 grant_type: "refresh_token",
931 client_id,
932 refresh_token,
933 }
934 }
935}
936
937#[derive(Debug, Serialize, Deserialize)]
938pub(super) struct ClientCredentialsRequest {
939 grant_type: &'static str,
940 scope: Option<&'static str>,
941}
942
943impl ClientCredentialsRequest {
944 pub(super) const fn new(scope: Option<&'static str>) -> Self {
945 Self {
946 grant_type: "client_credentials",
947 scope,
948 }
949 }
950}
951
952#[derive(Deserialize, Debug, Serialize)]
953pub(super) struct RefreshTokenResponse {
954 pub(super) refresh_token: Option<SecretRefreshToken>,
955 pub(super) access_token: SecretAccessToken,
956}
957
958#[async_trait::async_trait]
960pub trait TokenRefresher: Clone + std::fmt::Debug + Send {
961 type Error;
964
965 async fn validated_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
967
968 async fn get_access_token(&self) -> Result<Option<SecretAccessToken>, Self::Error>;
970
971 async fn refresh_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
973
974 #[cfg(feature = "tracing")]
976 fn base_url(&self) -> &str;
977
978 #[cfg(feature = "tracing-config")]
980 fn tracing_configuration(&self) -> Option<&TracingConfiguration>;
981
982 #[cfg(feature = "tracing")]
985 #[allow(clippy::needless_return)]
986 fn should_trace(&self, url: &UrlPatternMatchInput) -> bool {
987 #[cfg(not(feature = "tracing-config"))]
988 {
989 let _ = url;
990 return true;
991 }
992
993 #[cfg(feature = "tracing-config")]
994 self.tracing_configuration()
995 .is_none_or(|config| config.is_enabled(url))
996 }
997}
998
999#[async_trait::async_trait]
1000impl TokenRefresher for ClientConfiguration {
1001 type Error = TokenError;
1002
1003 async fn validated_access_token(&self) -> Result<SecretAccessToken, Self::Error> {
1004 self.get_bearer_access_token().await
1005 }
1006
1007 async fn refresh_access_token(&self) -> Result<SecretAccessToken, Self::Error> {
1008 match self.refresh().await {
1009 Ok(session) => Ok(session.access_token()?.clone()),
1010 Err(TokenError::Write {
1011 error: _error,
1012 oauth_session,
1013 }) => {
1014 #[cfg(feature = "tracing")]
1016 tracing::warn!(
1017 "Token refresh succeeded but failed to persist: {_error}. Returning access token from error.",
1018 );
1019 Ok(oauth_session.access_token()?.clone())
1020 }
1021 Err(e) => Err(e),
1022 }
1023 }
1024
1025 async fn get_access_token(&self) -> Result<Option<SecretAccessToken>, Self::Error> {
1026 Ok(Some(self.oauth_session().await?.access_token()?.clone()))
1027 }
1028
1029 #[cfg(feature = "tracing")]
1030 fn base_url(&self) -> &str {
1031 &self.grpc_api_url
1032 }
1033
1034 #[cfg(feature = "tracing-config")]
1035 fn tracing_configuration(&self) -> Option<&TracingConfiguration> {
1036 self.tracing_configuration.as_ref()
1037 }
1038}
1039
1040pub(super) fn default_http_client()
1042-> Result<qcs_dependencies_client::reqwest::Client, qcs_dependencies_client::reqwest::Error> {
1043 qcs_dependencies_client::reqwest::Client::builder()
1044 .timeout(std::time::Duration::from_secs(10))
1045 .build()
1046}
1047
1048#[cfg(test)]
1049mod test {
1050 #![allow(clippy::result_large_err, reason = "happens in figment tests")]
1051
1052 use std::time::Duration;
1053
1054 use super::*;
1055 use httpmock::prelude::*;
1056 use rstest::rstest;
1057 use time::format_description::well_known::Rfc3339;
1058 use tokio::time::Instant;
1059 use toml_edit::DocumentMut;
1060
1061 #[tokio::test]
1062 async fn test_tokens_blocked_during_refresh() {
1063 let mock_server = MockServer::start_async().await;
1064
1065 let oidc_mock = mock_server
1066 .mock_async(|when, then| {
1067 when.method(GET).path("/.well-known/openid-configuration");
1068 then.status(200)
1069 .json_body_obj(&oidc::Discovery::new_for_test(
1070 mock_server.base_url().parse().unwrap(),
1071 ));
1072 })
1073 .await;
1074
1075 let issuer_mock = mock_server
1076 .mock_async(|when, then| {
1077 when.method(POST).path("/v1/token");
1078
1079 then.status(200)
1080 .delay(Duration::from_secs(3))
1081 .json_body_obj(&RefreshTokenResponse {
1082 access_token: SecretAccessToken::from("new_access"),
1083 refresh_token: Some(SecretRefreshToken::from("new_refresh")),
1084 });
1085 })
1086 .await;
1087
1088 let original_tokens = OAuthSession::from_refresh_token(
1089 RefreshToken::new(SecretRefreshToken::from("refresh")),
1090 AuthServer {
1091 client_id: "client_id".to_string(),
1092 issuer: mock_server.base_url(),
1093 scopes: None,
1094 },
1095 None,
1096 );
1097 let dispatcher: TokenDispatcher = original_tokens.clone().into();
1098 let dispatcher_clone1 = dispatcher.clone();
1099 let dispatcher_clone2 = dispatcher.clone();
1100
1101 let refresh_duration = Duration::from_secs(3);
1102
1103 let start_write = Instant::now();
1104 let write_future = tokio::spawn(async move {
1105 dispatcher_clone1
1106 .refresh(&ConfigSource::Default, "")
1107 .await
1108 .unwrap()
1109 });
1110
1111 let start_read = Instant::now();
1112 let read_future = tokio::spawn(async move { dispatcher_clone2.tokens().await });
1113
1114 let _ = write_future.await.unwrap();
1115 let read_result = read_future.await.unwrap();
1116
1117 let write_duration = start_write.elapsed();
1118 let read_duration = start_read.elapsed();
1119
1120 oidc_mock.assert_async().await;
1121 issuer_mock.assert_async().await;
1122
1123 assert!(
1124 write_duration >= refresh_duration,
1125 "Write operation did not take enough time"
1126 );
1127 assert!(
1128 read_duration >= refresh_duration,
1129 "Read operation was not blocked by the write operation"
1130 );
1131 assert_eq!(
1132 read_result.access_token.unwrap(),
1133 SecretAccessToken::from("new_access")
1134 );
1135 if let OAuthGrant::RefreshToken(payload) = read_result.payload {
1136 assert_eq!(
1137 payload.refresh_token,
1138 SecretRefreshToken::from("new_refresh")
1139 );
1140 } else {
1141 panic!(
1142 "Expected RefreshToken payload, got {:?}",
1143 read_result.payload
1144 );
1145 }
1146 }
1147
1148 #[tokio::test]
1152 async fn test_refresh_token_request_rejected_by_auth_server() {
1153 let mock_server = MockServer::start_async().await;
1154
1155 let oidc_mock = mock_server
1156 .mock_async(|when, then| {
1157 when.method(GET).path("/.well-known/openid-configuration");
1158 then.status(200)
1159 .json_body_obj(&oidc::Discovery::new_for_test(
1160 mock_server.base_url().parse().unwrap(),
1161 ));
1162 })
1163 .await;
1164
1165 let issuer_mock = mock_server
1166 .mock_async(|when, then| {
1167 when.method(POST).path("/v1/token");
1168 then.status(400).json_body_obj(&serde_json::json!({
1169 "error": "invalid_grant",
1170 "error_description": "Unknown or invalid refresh token.",
1171 }));
1172 })
1173 .await;
1174
1175 let mut refresh_token = RefreshToken::new(SecretRefreshToken::from("revoked_refresh"));
1176 let auth_server = AuthServer {
1177 client_id: "client_id".to_string(),
1178 issuer: mock_server.base_url(),
1179 scopes: None,
1180 };
1181
1182 let result = refresh_token.request_access_token(&auth_server).await;
1183
1184 oidc_mock.assert_async().await;
1185 issuer_mock.assert_async().await;
1186
1187 assert!(
1188 result.is_err(),
1189 "a rejected refresh token request should be an error, got {result:?}"
1190 );
1191 }
1192
1193 #[rstest]
1194 fn test_qcs_secrets_readonly(
1195 #[values(
1196 (Some("TRUE"), true),
1197 (Some("tRue"), true),
1198 (Some("true"), true),
1199 (Some("YES"), true),
1200 (Some("yEs"), true),
1201 (Some("yes"), true),
1202 (Some("1"), true),
1203 (Some("2"), false),
1204 (Some("other"), false),
1205 (Some(""), false),
1206 (None, false),
1207 )]
1208 read_only_values: (Option<&str>, bool),
1209 #[values(true, false)] read_only_perm: bool,
1210 ) {
1211 let (maybe_read_only_env, env_is_read_only) = read_only_values;
1212 let expected_update = !env_is_read_only && !read_only_perm;
1213 figment::Jail::expect_with(|jail| {
1214 let profile_name = "test";
1215 let initial_access_token = "initial_access_token";
1216 let initial_refresh_token = "initial_refresh_token";
1217
1218 let initial_secrets_file_contents = format!(
1219 r#"
1220[credentials]
1221[credentials.{profile_name}]
1222[credentials.{profile_name}.token_payload]
1223access_token = "{initial_access_token}"
1224expires_in = 3600
1225id_token = "id_token"
1226refresh_token = "{initial_refresh_token}"
1227scope = "offline_access openid profile email"
1228token_type = "Bearer"
1229updated_at = "2024-01-01T00:00:00Z"
1230"#
1231 );
1232
1233 jail.clear_env();
1235
1236 let secrets_path = "secrets.toml";
1238 jail.create_file(secrets_path, initial_secrets_file_contents.as_str())
1239 .expect("should create test secrets.toml");
1240
1241 if read_only_perm {
1242 let mut permissions = std::fs::metadata(secrets_path)
1243 .expect("Should be able to get file metadata")
1244 .permissions();
1245 permissions.set_readonly(true);
1246 std::fs::set_permissions(secrets_path, permissions)
1247 .expect("Should be able to set file permissions");
1248 }
1249
1250 let rt = tokio::runtime::Runtime::new().unwrap();
1251 rt.block_on(async {
1252 let mock_server = MockServer::start_async().await;
1253
1254 let oidc_mock = mock_server
1255 .mock_async(|when, then| {
1256 when.method(GET).path("/.well-known/openid-configuration");
1257 then.status(200)
1258 .json_body_obj(&oidc::Discovery::new_for_test(mock_server.base_url().parse().unwrap()));
1259 })
1260 .await;
1261
1262 let new_access_token = SecretAccessToken::from("new_access_token");
1264 let issuer_mock = mock_server
1265 .mock_async(|when, then| {
1266 when.method(POST).path("/v1/token");
1267 then.status(200).json_body_obj(&RefreshTokenResponse {
1268 access_token: new_access_token.clone(),
1269 refresh_token: Some(SecretRefreshToken::from(initial_refresh_token)),
1270 });
1271 })
1272 .await;
1273
1274 let original_tokens = OAuthSession::from_refresh_token(
1276 RefreshToken::new(SecretRefreshToken::from(initial_refresh_token)),
1277 AuthServer { client_id: "client_id".to_string(), issuer: mock_server.base_url(), scopes: None },
1278 Some(SecretAccessToken::from(initial_refresh_token)),
1279 );
1280 let dispatcher: TokenDispatcher = original_tokens.into();
1281
1282 jail.set_env("QCS_SECRETS_FILE_PATH", "secrets.toml");
1284 jail.set_env("QCS_PROFILE_NAME", "test");
1285 if let Some(read_only_env) = maybe_read_only_env {
1286 jail.set_env("QCS_SECRETS_READ_ONLY", read_only_env);
1287 }
1288
1289 let before_refresh = OffsetDateTime::now_utc();
1290
1291 dispatcher
1292 .refresh(
1293 &ConfigSource::File {
1294 settings_path: "".into(),
1295 secrets_path: "secrets.toml".into(),
1296 },
1297 profile_name,
1298 )
1299 .await
1300 .unwrap();
1301
1302 oidc_mock.assert_async().await;
1303 issuer_mock.assert_async().await;
1304
1305 let content = std::fs::read_to_string("secrets.toml").unwrap();
1307 if !expected_update {
1308 assert!(
1309 content.eq(initial_secrets_file_contents.as_str()),
1310 "File should not be updated when QCS_SECRETS_READ_ONLY is set or file permissions are read-only"
1311 );
1312 return;
1313 }
1314
1315 let mut toml = std::fs::read_to_string(secrets_path)
1317 .unwrap()
1318 .parse::<DocumentMut>()
1319 .unwrap();
1320
1321 let token_payload = toml
1322 .get_mut("credentials")
1323 .and_then(|credentials| {
1324 credentials.get_mut(profile_name)?.get_mut("token_payload")
1325 })
1326 .expect("Should be able to get token_payload table");
1327
1328 let access_token = token_payload.get("access_token").unwrap().as_str().map(str::to_string).map(SecretAccessToken::from);
1329
1330 assert_eq!(
1331 access_token,
1332 Some(new_access_token)
1333 );
1334
1335 assert!(
1336 OffsetDateTime::parse(
1337 token_payload.get("updated_at").unwrap().as_str().unwrap(),
1338 &Rfc3339
1339 )
1340 .unwrap()
1341 > before_refresh
1342 );
1343
1344 let content = std::fs::read_to_string("secrets.toml").unwrap();
1345 assert!(
1346 content.contains("new_access_token"),
1347 "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"
1348 );
1349 });
1350 Ok(())
1351 });
1352 }
1353
1354 #[test]
1357 fn test_refresh_token_grant_persists_rotated_refresh_token() {
1358 let initial_refresh_token = "initial_refresh_token";
1359 let rotated_refresh_token = "rotated_refresh_token";
1360 let new_access_token = "new_access_token";
1361
1362 figment::Jail::expect_with(|jail| {
1363 jail.clear_env();
1364
1365 let secrets_path = "secrets.toml";
1366 let initial_secrets_file_contents = format!(
1367 r#"
1368[credentials]
1369[credentials.test]
1370[credentials.test.token_payload]
1371access_token = "initial_access_token"
1372refresh_token = "{initial_refresh_token}"
1373updated_at = "2024-01-01T00:00:00Z"
1374"#
1375 );
1376 jail.create_file(secrets_path, &initial_secrets_file_contents)
1377 .expect("should create test secrets.toml");
1378
1379 let rt = tokio::runtime::Runtime::new().unwrap();
1380 rt.block_on(async {
1381 let mock_server = MockServer::start_async().await;
1382 let oidc_mock = mock_server
1383 .mock_async(|when, then| {
1384 when.method(GET).path("/.well-known/openid-configuration");
1385 then.status(200)
1386 .json_body_obj(&oidc::Discovery::new_for_test(
1387 mock_server.base_url().parse().unwrap(),
1388 ));
1389 })
1390 .await;
1391 let issuer_mock = mock_server
1392 .mock_async(|when, then| {
1393 when.method(POST).path("/v1/token");
1394 then.status(200).json_body_obj(&RefreshTokenResponse {
1395 access_token: SecretAccessToken::from(new_access_token),
1396 refresh_token: Some(SecretRefreshToken::from(rotated_refresh_token)),
1397 });
1398 })
1399 .await;
1400
1401 let dispatcher: TokenDispatcher = OAuthSession::from_refresh_token(
1402 RefreshToken::new(SecretRefreshToken::from(initial_refresh_token)),
1403 AuthServer {
1404 client_id: "client_id".to_string(),
1405 issuer: mock_server.base_url(),
1406 scopes: None,
1407 },
1408 Some(SecretAccessToken::from("initial_access_token")),
1409 )
1410 .into();
1411
1412 dispatcher
1413 .refresh(
1414 &ConfigSource::File {
1415 settings_path: "".into(),
1416 secrets_path: secrets_path.into(),
1417 },
1418 "test",
1419 )
1420 .await
1421 .expect("refresh should succeed");
1422
1423 oidc_mock.assert_async().await;
1424 issuer_mock.assert_async().await;
1425 });
1426
1427 let Credential::TokenPayload(payload) = Secrets::load_from_path(&secrets_path.into())
1429 .expect("should load secrets")
1430 .credentials
1431 .remove("test")
1432 .expect("should have test credentials")
1433 else {
1434 panic!("expected a token payload credential");
1435 };
1436 assert_eq!(
1437 payload.refresh_token.unwrap(),
1438 SecretRefreshToken::from(rotated_refresh_token),
1439 "rotated refresh token should be persisted to the secrets file"
1440 );
1441 assert_eq!(
1442 payload.access_token.unwrap(),
1443 SecretAccessToken::from(new_access_token),
1444 "new access token should be persisted to the secrets file"
1445 );
1446
1447 Ok(())
1448 });
1449 }
1450
1451 #[test]
1452 fn test_auth_session_debug_fmt() {
1453 let session = OAuthSession {
1454 payload: OAuthGrant::ClientCredentials(ClientCredentials::new(
1455 "hidden_id",
1456 "hidden_secret",
1457 )),
1458 access_token: Some(SecretAccessToken::from("token")),
1459 auth_server: AuthServer {
1460 client_id: "some_id".into(),
1461 issuer: "some_url".into(),
1462 scopes: None,
1463 },
1464 };
1465
1466 assert_eq!(
1467 "OAuthSession { payload: ClientCredentials, access_token: Some(()), auth_server: AuthServer { client_id: \"some_id\", issuer: \"some_url\", scopes: None } }",
1468 &format!("{session:?}")
1469 );
1470 }
1471}