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 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#[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 pub refresh_token: SecretRefreshToken,
48}
49
50impl RefreshToken {
51 #[must_use]
53 pub const fn new(refresh_token: SecretRefreshToken) -> Self {
54 Self { refresh_token }
55 }
56
57 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 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#[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 pub client_id: String,
130 pub client_secret: ClientSecret,
132}
133
134impl ClientCredentials {
135 #[must_use]
136 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 #[must_use]
146 pub fn client_id(&self) -> &str {
147 &self.client_id
148 }
149
150 #[must_use]
152 pub const fn client_secret(&self) -> &ClientSecret {
153 &self.client_secret
154 }
155
156 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)]
196pub struct AuthTokens {
201 pub access_token: SecretAccessToken,
203 pub refresh_token: Option<RefreshToken>,
205}
206
207#[derive(Debug, thiserror::Error)]
209#[non_exhaustive]
210pub enum LoginError {
211 #[error(transparent)]
213 Pkce(#[from] PkceLoginError),
214 #[error(transparent)]
216 Device(#[from] DeviceLoginError),
217 #[error(transparent)]
219 Discovery(#[from] DiscoveryError),
220 #[error(transparent)]
222 Request(#[from] qcs_dependencies_client::reqwest::Error),
223}
224
225pub const LOGIN_FLOW_VAR: &str = "QCS_LOGIN_FLOW";
228
229#[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 #[default]
238 Auto,
239 Device,
241 Pkce,
243}
244
245#[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 #[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#[derive(Debug, PartialEq, Eq)]
289enum LoginFlow {
290 Pkce,
292 Device(url::Url),
294 DeviceThenPkce(url::Url),
296}
297
298impl LoginFlow {
299 fn select(
306 preference: LoginFlowPreference,
307 device_authorization_endpoint: Option<url::Url>,
308 ) -> Result<Self, DeviceLoginError> {
309 match (preference, device_authorization_endpoint) {
310 (LoginFlowPreference::Device, None) => Err(DeviceLoginError::NotSupported),
312 (LoginFlowPreference::Pkce, _) | (LoginFlowPreference::Auto, None) => Ok(Self::Pkce),
314 (LoginFlowPreference::Device, Some(endpoint)) => Ok(Self::Device(endpoint)),
316 (LoginFlowPreference::Auto, Some(endpoint)) => Ok(Self::DeviceThenPkce(endpoint)),
317 }
318 }
319}
320
321pub(crate) struct LoginFlowOptions {
323 pub(crate) preference: LoginFlowPreference,
325 pub(crate) redirect: RedirectBinding,
327}
328
329impl LoginFlowOptions {
330 pub(crate) fn from_env() -> Self {
333 Self::with_preference(LoginFlowPreference::from_env())
334 }
335
336 pub(crate) fn with_preference(preference: LoginFlowPreference) -> Self {
338 Self {
339 preference,
340 redirect: RedirectBinding::default(),
341 }
342 }
343}
344
345impl AuthTokens {
346 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 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 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 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
460async 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
480async 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))]
513pub enum OAuthGrant {
516 RefreshToken(RefreshToken),
518 ClientCredentials(ClientCredentials),
520 ExternallyManaged(ExternallyManaged),
522 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 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#[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 payload: OAuthGrant,
606 access_token: Option<SecretAccessToken>,
608 auth_server: AuthServer,
610}
611
612impl OAuthSession {
613 #[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 #[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 #[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 #[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 #[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 pub fn access_token(&self) -> Result<&SecretAccessToken, TokenError> {
703 self.access_token.as_ref().ok_or(TokenError::NoAccessToken)
704 }
705
706 #[must_use]
708 pub const fn payload(&self) -> &OAuthGrant {
709 &self.payload
710 }
711
712 #[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 #[must_use]
725 pub const fn auth_server(&self) -> &AuthServer {
726 &self.auth_server
727 }
728
729 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
743pub(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
776pub(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 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 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#[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 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 #[must_use]
888 pub async fn tokens(&self) -> OAuthSession {
889 self.use_tokens(Clone::clone).await
890 }
891
892 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 pub async fn validate(&self) -> Result<SecretAccessToken, TokenError> {
914 self.use_tokens(OAuthSession::validate).await
915 }
916
917 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 *self.refreshing.lock().await = false;
946 self.notify_refreshed.notify_waiters();
947
948 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 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
975pub type RefreshFunction = Box<dyn (Fn(AuthServer) -> RefreshResult) + Send + Sync>;
977
978#[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 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 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 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 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#[async_trait::async_trait]
1187pub trait TokenRefresher: Clone + std::fmt::Debug + Send {
1188 type Error;
1191
1192 async fn validated_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
1194
1195 async fn get_access_token(&self) -> Result<Option<SecretAccessToken>, Self::Error>;
1197
1198 async fn refresh_access_token(&self) -> Result<SecretAccessToken, Self::Error>;
1200
1201 #[cfg(feature = "tracing")]
1203 fn base_url(&self) -> &str;
1204
1205 #[cfg(feature = "tracing-config")]
1207 fn tracing_configuration(&self) -> Option<&TracingConfiguration>;
1208
1209 #[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 #[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
1267pub 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 #[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 jail.clear_env();
1468
1469 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 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 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 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 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 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 #[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 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 #[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 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 assert!(matches!(
1736 select(LoginFlowPreference::Device, false),
1737 Err(DeviceLoginError::NotSupported)
1738 ));
1739 }
1740
1741 #[tokio::test(flavor = "multi_thread")]
1745 async fn test_device_flow_falls_back_to_pkce_when_rejected() {
1746 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 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}