diff --git a/bindings/matrix-sdk-ffi/src/authentication.rs b/bindings/matrix-sdk-ffi/src/authentication.rs index 2822fe6b5..028b16b1f 100644 --- a/bindings/matrix-sdk-ffi/src/authentication.rs +++ b/bindings/matrix-sdk-ffi/src/authentication.rs @@ -6,6 +6,7 @@ use std::{ use matrix_sdk::{ authentication::oidc::{ + error::OauthAuthorizationCodeError, registrations::OidcRegistrationsError, types::{ iana::oauth::OAuthClientAuthenticationMethod, @@ -198,9 +199,13 @@ impl From for OidcError { match e { SdkOidcError::Discovery(error) if error.is_not_supported() => OidcError::NotSupported, SdkOidcError::MissingRedirectUri => OidcError::MetadataInvalid, - SdkOidcError::InvalidCallbackUrl => OidcError::CallbackUrlInvalid, - SdkOidcError::InvalidState => OidcError::CallbackUrlInvalid, - SdkOidcError::CancelledAuthorization => OidcError::Cancelled, + SdkOidcError::AuthorizationCode(OauthAuthorizationCodeError::RedirectUri(_)) + | SdkOidcError::AuthorizationCode(OauthAuthorizationCodeError::InvalidState) => { + OidcError::CallbackUrlInvalid + } + SdkOidcError::AuthorizationCode(OauthAuthorizationCodeError::Cancelled) => { + OidcError::Cancelled + } _ => OidcError::Generic { message: e.to_string() }, } } diff --git a/bindings/matrix-sdk-ffi/src/client.rs b/bindings/matrix-sdk-ffi/src/client.rs index 12a2a7caa..ad3dad869 100644 --- a/bindings/matrix-sdk-ffi/src/client.rs +++ b/bindings/matrix-sdk-ffi/src/client.rs @@ -20,6 +20,7 @@ use matrix_sdk::{ }, ruma::{ api::client::{ + discovery::get_authorization_server_metadata::msc2965::Prompt as RumaOidcPrompt, push::{EmailPusherData, PusherIds, PusherInit, PusherKind as RumaPusherKind}, room::{create_room, Visibility}, session::get_login_types, @@ -399,7 +400,7 @@ impl Client { pub async fn url_for_oidc( &self, oidc_configuration: &OidcConfiguration, - prompt: OidcPrompt, + prompt: Option, ) -> Result, OidcError> { let oidc_metadata: VerifiedClientMetadata = oidc_configuration.try_into()?; let registrations_file = Path::new(&oidc_configuration.dynamic_registrations_file); @@ -420,8 +421,11 @@ impl Client { static_registrations, )?; - let data = - self.inner.oidc().url_for_oidc(oidc_metadata, registrations, prompt.into()).await?; + let data = self + .inner + .oidc() + .url_for_oidc(oidc_metadata, registrations, prompt.map(Into::into)) + .await?; Ok(Arc::new(data)) } @@ -1813,26 +1817,6 @@ impl TryFrom for SdkSlidingSyncVersion { #[derive(Clone, uniffi::Enum)] pub enum OidcPrompt { - /// The Authorization Server must not display any authentication or consent - /// user interface pages. - None, - - /// The Authorization Server should prompt the End-User for - /// reauthentication. - Login, - - /// The Authorization Server should prompt the End-User for consent before - /// returning information to the Client. - Consent, - - /// The Authorization Server should prompt the End-User to select a user - /// account. - /// - /// This enables an End-User who has multiple accounts at the Authorization - /// Server to select amongst the multiple accounts that they might have - /// current sessions for. - SelectAccount, - /// The Authorization Server should prompt the End-User to create a user /// account. /// @@ -1846,26 +1830,17 @@ pub enum OidcPrompt { impl From<&SdkOidcPrompt> for OidcPrompt { fn from(value: &SdkOidcPrompt) -> Self { match value { - SdkOidcPrompt::None => Self::None, - SdkOidcPrompt::Login => Self::Login, - SdkOidcPrompt::Consent => Self::Consent, - SdkOidcPrompt::SelectAccount => Self::SelectAccount, SdkOidcPrompt::Create => Self::Create, - SdkOidcPrompt::Unknown(value) => Self::Unknown { value: value.to_owned() }, _ => Self::Unknown { value: value.to_string() }, } } } -impl From for SdkOidcPrompt { +impl From for RumaOidcPrompt { fn from(value: OidcPrompt) -> Self { match value { - OidcPrompt::None => Self::None, - OidcPrompt::Login => Self::Login, - OidcPrompt::Consent => Self::Consent, - OidcPrompt::SelectAccount => Self::SelectAccount, OidcPrompt::Create => Self::Create, - OidcPrompt::Unknown { value } => Self::Unknown(value), + OidcPrompt::Unknown { value } => value.into(), } } } diff --git a/crates/matrix-sdk/Cargo.toml b/crates/matrix-sdk/Cargo.toml index 676c66bd0..0519d137a 100644 --- a/crates/matrix-sdk/Cargo.toml +++ b/crates/matrix-sdk/Cargo.toml @@ -93,7 +93,7 @@ matrix-sdk-sqlite = { workspace = true, optional = true } matrix-sdk-test = { workspace = true, optional = true } mime = { workspace = true } mime2ext = "0.1.53" -oauth2 = { version = "5.0.0", default-features = false, features = ["reqwest"], optional = true } +oauth2 = { version = "5.0.0", default-features = false, features = ["reqwest", "timing-resistant-secret-traits"], optional = true } once_cell = { workspace = true } percent-encoding = "2.3.1" pin-project-lite = { workspace = true } diff --git a/crates/matrix-sdk/src/authentication/oidc/auth_code_builder.rs b/crates/matrix-sdk/src/authentication/oidc/auth_code_builder.rs index 1b8701727..84df54bbb 100644 --- a/crates/matrix-sdk/src/authentication/oidc/auth_code_builder.rs +++ b/crates/matrix-sdk/src/authentication/oidc/auth_code_builder.rs @@ -12,16 +12,17 @@ // See the License for the specific language governing permissions and // limitations under the License. -use mas_oidc_client::{ - requests::authorization_code::{build_authorization_url, AuthorizationRequestData}, - types::{requests::Prompt, scope::Scope}, +use std::borrow::Cow; + +use oauth2::{ + basic::BasicClient as OauthClient, AuthUrl, CsrfToken, PkceCodeChallenge, RedirectUrl, Scope, }; -use ruma::UserId; +use ruma::{api::client::discovery::get_authorization_server_metadata::msc2965::Prompt, UserId}; use tracing::{info, instrument}; use url::Url; use super::{Oidc, OidcError}; -use crate::Result; +use crate::{authentication::oidc::AuthorizationValidationData, Result}; /// Builder type used to configure optional settings for authorization with an /// OpenID Connect Provider via the Authorization Code flow. @@ -30,15 +31,15 @@ use crate::Result; #[allow(missing_debug_implementations)] pub struct OidcAuthCodeUrlBuilder { oidc: Oidc, - scope: Scope, + scopes: Vec, redirect_uri: Url, prompt: Option>, login_hint: Option, } impl OidcAuthCodeUrlBuilder { - pub(super) fn new(oidc: Oidc, scope: Scope, redirect_uri: Url) -> Self { - Self { oidc, scope, redirect_uri, prompt: None, login_hint: None } + pub(super) fn new(oidc: Oidc, scopes: Vec, redirect_uri: Url) -> Self { + Self { oidc, scopes, redirect_uri, prompt: None, login_hint: None } } /// Set the [`Prompt`] of the authorization URL. @@ -73,34 +74,44 @@ impl OidcAuthCodeUrlBuilder { /// request fails. #[instrument(target = "matrix_sdk::client", skip_all)] pub async fn build(self) -> Result { - let Self { oidc, scope, redirect_uri, prompt, login_hint } = self; + let Self { oidc, scopes, redirect_uri, prompt, login_hint } = self; let data = oidc.data().ok_or(OidcError::NotAuthenticated)?; info!( issuer = data.issuer, - %scope, "Authorizing scope via the OpenID Connect Authorization Code flow" + ?scopes, + "Authorizing scope via the OpenID Connect Authorization Code flow" ); let provider_metadata = oidc.provider_metadata().await?; + let auth_url = AuthUrl::from_url(provider_metadata.authorization_endpoint().clone()); - let mut authorization_data = - AuthorizationRequestData::new(data.client_id.as_str().to_owned(), scope, redirect_uri); - authorization_data.code_challenge_methods_supported = - provider_metadata.code_challenge_methods_supported.clone(); - authorization_data.prompt = prompt; - authorization_data.login_hint = login_hint; + let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256(); + let redirect_uri = RedirectUrl::from_url(redirect_uri); - let authorization_endpoint = provider_metadata.authorization_endpoint(); + let client = OauthClient::new(data.client_id.clone()).set_auth_uri(auth_url); + let mut request = client + .authorize_url(CsrfToken::new_random) + .add_scopes(scopes) + .set_pkce_challenge(pkce_challenge) + .set_redirect_uri(Cow::Borrowed(&redirect_uri)); - let (url, validation_data) = build_authorization_url( - authorization_endpoint.clone(), - authorization_data, - &mut super::rng()?, - )?; + if let Some(prompt) = prompt { + // This should be a list of space separated values. + let prompt_str = prompt.iter().map(Prompt::as_str).collect::>().join(" "); + request = request.add_extra_param("prompt", prompt_str); + } - let state = validation_data.state.clone(); + if let Some(login_hint) = login_hint { + request = request.add_extra_param("login_hint", login_hint); + } - data.authorization_data.lock().await.insert(state.clone(), validation_data); + let (url, state) = request.url(); + + data.authorization_data + .lock() + .await + .insert(state.clone(), AuthorizationValidationData { redirect_uri, pkce_verifier }); Ok(OidcAuthorizationData { url, state }) } @@ -114,7 +125,7 @@ pub struct OidcAuthorizationData { pub url: Url, /// A unique identifier for the request, used to ensure the response /// originated from the authentication issuer. - pub state: String, + pub state: CsrfToken, } #[cfg(feature = "uniffi")] diff --git a/crates/matrix-sdk/src/authentication/oidc/error.rs b/crates/matrix-sdk/src/authentication/oidc/error.rs index e57ba6da2..eb5e4fe76 100644 --- a/crates/matrix-sdk/src/authentication/oidc/error.rs +++ b/crates/matrix-sdk/src/authentication/oidc/error.rs @@ -15,10 +15,13 @@ //! Error types used in the [`Oidc`](super::Oidc) API. pub use mas_oidc_client::error::*; +use matrix_sdk_base::deserialized_responses::PrivOwnedStr; +use oauth2::ErrorResponseType; pub use oauth2::{ basic::{BasicErrorResponse, BasicErrorResponseType, BasicRequestTokenError}, HttpClientError, RequestTokenError, StandardErrorResponse, }; +use ruma::serde::{PartialEqAsRefStr, StringEnum}; pub use super::cross_process::CrossProcessRefreshLockError; @@ -69,26 +72,9 @@ pub enum OidcError { #[error("client not authenticated")] NotAuthenticated, - /// The state used to complete authorization doesn't match an original - /// value. - #[error("the supplied state is unexpected")] - InvalidState, - - /// The user cancelled authorization in the web view. - #[error("authorization cancelled")] - CancelledAuthorization, - - /// The login was completed with an invalid callback. - #[error("the supplied callback URL is invalid")] - InvalidCallbackUrl, - - /// An error occurred during authorization. - #[error("authorization failed")] - Authorization(super::AuthorizationError), - - /// The device ID is invalid. - #[error("invalid device ID")] - InvalidDeviceId, + /// An error occurred using the OAuth 2.0 authorization code grant. + #[error("authorization code grant failed: {0}")] + AuthorizationCode(#[from] OauthAuthorizationCodeError), /// An error occurred interacting with the OAuth 2.0 authorization server /// while refreshing the access token. @@ -154,3 +140,87 @@ impl OauthDiscoveryError { matches!(self, Self::NotSupported) } } + +/// All errors that can occur when using the Authorization Code grant with the +/// OAuth 2.0 API. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum OauthAuthorizationCodeError { + /// The query of the redirect URI doesn't have the expected format. + #[error(transparent)] + RedirectUri(#[from] RedirectUriQueryParseError), + + /// The user cancelled the authorization in the web UI. + #[error("authorization cancelled by the user")] + Cancelled, + + /// An error occurred when getting the authorization from the user in the + /// web UI. + #[error("authorization failed: {0}")] + Authorization(StandardErrorResponse), + + /// The state used to complete authorization doesn't match any of the + /// ongoing authorizations. + #[error("authorization state value is unexpected")] + InvalidState, + + /// An error occurred interacting with the OAuth 2.0 authorization server + /// while exchanging the authorization code for an access token. + #[error("failed to request token: {0}")] + RequestToken(BasicRequestTokenError>), +} + +impl From> + for OauthAuthorizationCodeError +{ + fn from(value: StandardErrorResponse) -> Self { + if *value.error() == AuthorizationCodeErrorResponseType::AccessDenied { + // The user cancelled the login in the web view. + Self::Cancelled + } else { + Self::Authorization(value) + } + } +} + +/// Error response returned by server after requesting an authorization code. +/// +/// The fields in this structure are defined in [Section 4.1.2.1 of RFC 6749]. +/// +/// [Section 4.1.2.1 of RFC 6749]: https://datatracker.ietf.org/doc/html/rfc6749#section-4.1.2.1 +#[derive(Clone, StringEnum, PartialEqAsRefStr, Eq)] +#[ruma_enum(rename_all = "snake_case")] +#[non_exhaustive] +pub enum AuthorizationCodeErrorResponseType { + /// The request is invalid. + /// + /// It is missing a required parameter, includes an invalid parameter value, + /// includes a parameter more than once, or is otherwise malformed. + InvalidRequest, + + /// The client is not authorized to request an authorization code using this + /// method. + UnauthorizedClient, + + /// The resource owner or authorization server denied the request. + AccessDenied, + + /// The authorization server does not support obtaining an authorization + /// code using this method. + UnsupportedResponseType, + + /// The requested scope is invalid, unknown, or malformed. + InvalidScope, + + /// The authorization server encountered an unexpected error. + ServerError, + + /// The authorization server is currently unable to handle the request due + /// to a temporary overloading or maintenance of the server. + TemporarilyUnavailable, + + #[doc(hidden)] + _Custom(PrivOwnedStr), +} + +impl ErrorResponseType for AuthorizationCodeErrorResponseType {} diff --git a/crates/matrix-sdk/src/authentication/oidc/mod.rs b/crates/matrix-sdk/src/authentication/oidc/mod.rs index 846533e60..f0bbcb267 100644 --- a/crates/matrix-sdk/src/authentication/oidc/mod.rs +++ b/crates/matrix-sdk/src/authentication/oidc/mod.rs @@ -146,45 +146,52 @@ //! [`AuthenticateError::InsufficientScope`]: ruma::api::client::error::AuthenticateError //! [`examples/oidc_cli`]: https://github.com/matrix-org/matrix-rust-sdk/tree/main/examples/oidc_cli -use std::{collections::HashMap, fmt, future::Future, pin::Pin, sync::Arc}; +use std::{borrow::Cow, collections::HashMap, fmt, future::Future, pin::Pin, sync::Arc}; use as_variant::as_variant; use chrono::Utc; -use error::{CrossProcessRefreshLockError, OauthDiscoveryError, RedirectUriQueryParseError}; +use error::{ + CrossProcessRefreshLockError, OauthAuthorizationCodeError, OauthDiscoveryError, + RedirectUriQueryParseError, +}; use eyeball::SharedObservable; use futures_core::Stream; use mas_oidc_client::{ http_service::HttpService, requests::{ account_management::{build_account_management_url, AccountManagementActionFull}, - authorization_code::{access_token_with_authorization_code, AuthorizationValidationData}, discovery::{discover, insecure_discover}, registration::register_client, revocation::revoke_token, }, types::{ client_credentials::ClientCredentials, - errors::{ClientError, ClientErrorCode::AccessDenied}, iana::oauth::OAuthTokenTypeHint, oidc::{ AccountManagementAction, ProviderMetadata, ProviderMetadataVerificationError, VerifiedProviderMetadata, }, registration::{ClientRegistrationResponse, VerifiedClientMetadata}, - requests::Prompt, - scope::{MatrixApiScopeToken, ScopeToken}, }, }; pub use mas_oidc_client::{requests, types}; #[cfg(feature = "e2e-encryption")] use matrix_sdk_base::crypto::types::qr_login::QrCodeData; use matrix_sdk_base::{once_cell::sync::OnceCell, SessionMeta}; +pub use oauth2::CsrfToken; use oauth2::{ - basic::BasicClient as OauthClient, AsyncHttpClient, HttpRequest, HttpResponse, RefreshToken, - TokenResponse, TokenUrl, + basic::BasicClient as OauthClient, AsyncHttpClient, HttpRequest, HttpResponse, + PkceCodeVerifier, RedirectUrl, RefreshToken, Scope, StandardErrorResponse, TokenResponse, + TokenUrl, +}; +use rand::{rngs::StdRng, SeedableRng}; +use ruma::{ + api::client::discovery::{ + get_authentication_issuer, + get_authorization_server_metadata::{self, msc2965::Prompt}, + }, + DeviceId, OwnedDeviceId, }; -use rand::{rngs::StdRng, Rng, SeedableRng}; -use ruma::api::client::discovery::{get_authentication_issuer, get_authorization_server_metadata}; use serde::{Deserialize, Serialize}; use sha2::Digest as _; use tokio::{spawn, sync::Mutex}; @@ -243,7 +250,7 @@ pub(crate) struct OidcAuthData { pub(crate) client_id: ClientId, pub(crate) tokens: OnceCell>, /// The data necessary to validate authorization responses. - pub(crate) authorization_data: Mutex>, + authorization_data: Mutex>, } impl OidcAuthData { @@ -435,7 +442,7 @@ impl Oidc { &self, client_metadata: VerifiedClientMetadata, registrations: OidcRegistrations, - prompt: Prompt, + prompt: Option, ) -> Result { let metadata = self.provider_metadata().await?; @@ -447,7 +454,11 @@ impl Oidc { self.configure(metadata.issuer().to_owned(), client_metadata, registrations).await?; let mut data_builder = self.login(redirect_url.clone(), None)?; - data_builder = data_builder.prompt(vec![prompt]); + + if let Some(prompt) = prompt { + data_builder = data_builder.prompt(vec![prompt]); + } + let data = data_builder.build().await?; Ok(data) @@ -462,16 +473,13 @@ impl Oidc { callback_url: Url, ) -> Result<()> { let response = AuthorizationResponse::parse_uri(&callback_url) - .or(Err(OidcError::InvalidCallbackUrl))?; + .map_err(OauthAuthorizationCodeError::from) + .map_err(OidcError::from)?; let code = match response { AuthorizationResponse::Success(code) => code, AuthorizationResponse::Error(err) => { - if err.error.error == AccessDenied { - // The user cancelled the login in the web view. - return Err(OidcError::CancelledAuthorization.into()); - } - return Err(OidcError::Authorization(err).into()); + return Err(OidcError::from(OauthAuthorizationCodeError::from(err.error)).into()); } }; @@ -479,7 +487,7 @@ impl Oidc { // the client to have called `abort_authorization` which we can't guarantee so // lets double check with their supplied authorization data to be safe. if code.state != authorization_data.state { - return Err(OidcError::InvalidState.into()); + return Err(OidcError::from(OauthAuthorizationCodeError::InvalidState).into()); }; self.finish_authorization(code).await?; @@ -1083,21 +1091,20 @@ impl Oidc { } /// The scopes to request for logging in. - fn login_scopes(device_id: Option) -> Result<[ScopeToken; 3], OidcError> { - // Generate the device ID if it is not provided. - let device_id = device_id.unwrap_or_else(|| { - rand::thread_rng() - .sample_iter(&rand::distributions::Alphanumeric) - .map(char::from) - .take(10) - .collect::() - }); + fn login_scopes(device_id: Option) -> [Scope; 2] { + /// Scope to grand full access to the client-server API. + const SCOPE_MATRIX_CLIENT_SERVER_API_FULL_ACCESS: &str = + "urn:matrix:org.matrix.msc2967.client:api:*"; + /// Prefix of the scope to bind a device ID to an access token. + const SCOPE_MATRIX_DEVICE_ID_PREFIX: &str = "urn:matrix:org.matrix.msc2967.client:device:"; - Ok([ - ScopeToken::Openid, - ScopeToken::MatrixApi(MatrixApiScopeToken::Full), - ScopeToken::try_with_matrix_device(device_id).or(Err(OidcError::InvalidDeviceId))?, - ]) + // Generate the device ID if it is not provided. + let device_id = device_id.unwrap_or_else(DeviceId::new); + + [ + Scope::new(SCOPE_MATRIX_CLIENT_SERVER_API_FULL_ACCESS.to_owned()), + Scope::new(format!("{SCOPE_MATRIX_DEVICE_ID_PREFIX}{device_id}")), + ] } /// Login via OpenID Connect with the Authorization Code flow. @@ -1174,11 +1181,11 @@ impl Oidc { pub fn login( &self, redirect_uri: Url, - device_id: Option, + device_id: Option, ) -> Result { - let scope = Self::login_scopes(device_id)?.into_iter().collect(); + let scopes = Self::login_scopes(device_id).to_vec(); - Ok(OidcAuthCodeUrlBuilder::new(self.clone(), scope, redirect_uri)) + Ok(OidcAuthCodeUrlBuilder::new(self.clone(), scopes, redirect_uri)) } /// Finish the login process. @@ -1265,30 +1272,30 @@ impl Oidc { auth_code: AuthorizationCode, ) -> Result<(), OidcError> { let data = self.data().ok_or(OidcError::NotAuthenticated)?; + let client_id = data.client_id.clone(); + let validation_data = data .authorization_data .lock() .await .remove(&auth_code.state) - .ok_or(OidcError::InvalidState)?; + .ok_or(OauthAuthorizationCodeError::InvalidState)?; let provider_metadata = self.provider_metadata().await?; + let token_uri = TokenUrl::from_url(provider_metadata.token_endpoint().clone()); - let (response, _) = access_token_with_authorization_code( - &self.http_service(), - data.credentials(), - provider_metadata.token_endpoint(), - auth_code.code, - validation_data, - None, - Utc::now(), - &mut rng()?, - ) - .await?; + let response = OauthClient::new(client_id) + .set_token_uri(token_uri) + .exchange_code(oauth2::AuthorizationCode::new(auth_code.code)) + .set_pkce_verifier(validation_data.pkce_verifier) + .set_redirect_uri(Cow::Owned(validation_data.redirect_uri)) + .request_async(self.http_client()) + .await + .map_err(OauthAuthorizationCodeError::RequestToken)?; self.set_session_tokens(OidcSessionTokens { - access_token: response.access_token, - refresh_token: response.refresh_token, + access_token: response.access_token().secret().clone(), + refresh_token: response.refresh_token().map(RefreshToken::secret).cloned(), }); Ok(()) @@ -1309,7 +1316,7 @@ impl Oidc { /// * `state` - The state received as part of the redirect URI when the /// authorization failed, or the one provided in [`OidcAuthorizationData`] /// after building the authorization URL. - pub async fn abort_authorization(&self, state: &str) { + pub async fn abort_authorization(&self, state: &CsrfToken) { if let Some(data) = self.data() { data.authorization_data.lock().await.remove(state); } @@ -1320,12 +1327,10 @@ impl Oidc { #[cfg(all(feature = "e2e-encryption", not(target_arch = "wasm32")))] async fn request_device_authorization( &self, - device_id: Option, + device_id: Option, ) -> Result { - let scopes = Self::login_scopes(device_id)? - .into_iter() - .map(|scope| oauth2::Scope::new(scope.to_string())); + let scopes = Self::login_scopes(device_id); let client_id = self.client_id().ok_or(OidcError::NotRegistered)?.clone(); @@ -1635,6 +1640,17 @@ impl fmt::Debug for OidcSessionTokens { } } +/// The data necessary to validate a response from the Token endpoint in the +/// Authorization Code flow. +#[derive(Debug)] +struct AuthorizationValidationData { + /// The URI where the end-user will be redirected after authorization. + redirect_uri: RedirectUrl, + + /// A string to correlate the authorization request to the token request. + pkce_verifier: PkceCodeVerifier, +} + /// The data returned by the provider in the redirect URI after a successful /// authorization. #[derive(Debug, Clone)] @@ -1680,7 +1696,7 @@ pub struct AuthorizationCode { /// The code to use to retrieve the access token. pub code: String, /// The unique identifier for this transaction. - pub state: String, + pub state: CsrfToken, } /// The data returned by the provider in the redirect URI after an authorization @@ -1689,9 +1705,9 @@ pub struct AuthorizationCode { pub struct AuthorizationError { /// The error. #[serde(flatten)] - pub error: ClientError, + pub error: StandardErrorResponse, /// The unique identifier for this transaction. - pub state: String, + pub state: CsrfToken, } fn rng() -> Result { diff --git a/crates/matrix-sdk/src/authentication/oidc/qrcode/login.rs b/crates/matrix-sdk/src/authentication/oidc/qrcode/login.rs index d350a6bbc..6fb47377b 100644 --- a/crates/matrix-sdk/src/authentication/oidc/qrcode/login.rs +++ b/crates/matrix-sdk/src/authentication/oidc/qrcode/login.rs @@ -294,7 +294,8 @@ impl<'a> LoginWithQrCode<'a> { device_id: Curve25519PublicKey, ) -> Result { let oidc = self.client.oidc(); - let response = oidc.request_device_authorization(Some(device_id.to_base64())).await?; + let response = + oidc.request_device_authorization(Some(device_id.to_base64().into())).await?; Ok(response) } diff --git a/crates/matrix-sdk/src/authentication/oidc/tests.rs b/crates/matrix-sdk/src/authentication/oidc/tests.rs index 96a709cb8..6baa6c599 100644 --- a/crates/matrix-sdk/src/authentication/oidc/tests.rs +++ b/crates/matrix-sdk/src/authentication/oidc/tests.rs @@ -3,14 +3,15 @@ use std::collections::HashMap; use anyhow::Context as _; use assert_matches::assert_matches; use mas_oidc_client::{ - requests::{ - account_management::AccountManagementActionFull, - authorization_code::AuthorizationValidationData, - }, - types::{errors::ClientErrorCode, registration::VerifiedClientMetadata, requests::Prompt}, + requests::account_management::AccountManagementActionFull, + types::registration::VerifiedClientMetadata, }; use matrix_sdk_test::async_test; -use ruma::ServerName; +use oauth2::{CsrfToken, PkceCodeChallenge, RedirectUrl}; +use ruma::{ + api::client::discovery::get_authorization_server_metadata::msc2965::Prompt, owned_device_id, + ServerName, +}; use serde_json::json; use stream_assert::{assert_next_matches, assert_pending}; use tempfile::tempdir; @@ -25,6 +26,10 @@ use super::{ Oidc, OidcError, OidcSessionTokens, RedirectUriQueryParseError, }; use crate::{ + authentication::oidc::{ + error::AuthorizationCodeErrorResponseType, AuthorizationValidationData, + OauthAuthorizationCodeError, + }, test_utils::{ client::{ oauth::{mock_client_metadata, mock_session, mock_session_tokens}, @@ -76,7 +81,7 @@ async fn test_high_level_login() -> anyhow::Result<()> { // When getting the OIDC login URL. let authorization_data = - oidc.url_for_oidc(metadata.clone(), registrations, Prompt::Login).await.unwrap(); + oidc.url_for_oidc(metadata.clone(), registrations, Some(Prompt::Create)).await.unwrap(); // Then the client should be configured correctly. assert!(oidc.issuer().is_some()); @@ -84,7 +89,7 @@ async fn test_high_level_login() -> anyhow::Result<()> { // When completing the login with a valid callback. let mut callback_uri = metadata.redirect_uris.clone().unwrap().first().unwrap().clone(); - callback_uri.set_query(Some(&format!("code=42&state={}", authorization_data.state))); + callback_uri.set_query(Some(&format!("code=42&state={}", authorization_data.state.secret()))); // Then the login should succeed. oidc.login_with_oidc_callback(&authorization_data, callback_uri).await?; @@ -97,20 +102,25 @@ async fn test_high_level_login_cancellation() -> anyhow::Result<()> { // Given a client ready to complete login. let (oidc, _server, metadata, registrations) = mock_environment().await.unwrap(); let authorization_data = - oidc.url_for_oidc(metadata.clone(), registrations, Prompt::Login).await.unwrap(); + oidc.url_for_oidc(metadata.clone(), registrations, None).await.unwrap(); assert!(oidc.issuer().is_some()); assert!(oidc.client_id().is_some()); // When completing login with a cancellation callback. let mut callback_uri = metadata.redirect_uris.clone().unwrap().first().unwrap().clone(); - callback_uri - .set_query(Some(&format!("error=access_denied&state={}", authorization_data.state))); + callback_uri.set_query(Some(&format!( + "error=access_denied&state={}", + authorization_data.state.secret() + ))); let error = oidc.login_with_oidc_callback(&authorization_data, callback_uri).await.unwrap_err(); // Then a cancellation error should be thrown. - assert_matches!(error, Error::Oidc(OidcError::CancelledAuthorization)); + assert_matches!( + error, + Error::Oidc(OidcError::AuthorizationCode(OauthAuthorizationCodeError::Cancelled)) + ); Ok(()) } @@ -120,7 +130,7 @@ async fn test_high_level_login_invalid_state() -> anyhow::Result<()> { // Given a client ready to complete login. let (oidc, _server, metadata, registrations) = mock_environment().await.unwrap(); let authorization_data = - oidc.url_for_oidc(metadata.clone(), registrations, Prompt::Login).await.unwrap(); + oidc.url_for_oidc(metadata.clone(), registrations, None).await.unwrap(); assert!(oidc.issuer().is_some()); assert!(oidc.client_id().is_some()); @@ -132,7 +142,10 @@ async fn test_high_level_login_invalid_state() -> anyhow::Result<()> { let error = oidc.login_with_oidc_callback(&authorization_data, callback_uri).await.unwrap_err(); // Then the login should fail by flagging the invalid state. - assert_matches!(error, Error::Oidc(OidcError::InvalidState)); + assert_matches!( + error, + Error::Oidc(OidcError::AuthorizationCode(OauthAuthorizationCodeError::InvalidState)) + ); Ok(()) } @@ -148,7 +161,7 @@ async fn test_login() -> anyhow::Result<()> { let client = server.client_builder().registered_with_oauth(server.server().uri()).build().await; let oidc = client.oidc(); - let device_id = "D3V1C31D".to_owned(); // yo this is 1999 speaking + let device_id = owned_device_id!("D3V1C31D"); // yo this is 1999 speaking let redirect_uri_str = REDIRECT_URI_STRING; let redirect_uri = Url::parse(redirect_uri_str)?; @@ -156,8 +169,8 @@ async fn test_login() -> anyhow::Result<()> { tracing::debug!("authorization data URL = {}", authorization_data.url); - let mut num_expected = 6; - let mut nonce = None; + let mut num_expected = 7; + let mut code_challenge = None; for (key, val) in authorization_data.url.query_pairs() { match &*key { @@ -174,16 +187,20 @@ async fn test_login() -> anyhow::Result<()> { num_expected -= 1; } "scope" => { - assert_eq!(val, format!("openid urn:matrix:org.matrix.msc2967.client:api:* urn:matrix:org.matrix.msc2967.client:device:{device_id}")); + assert_eq!(val, format!("urn:matrix:org.matrix.msc2967.client:api:* urn:matrix:org.matrix.msc2967.client:device:{device_id}")); num_expected -= 1; } "state" => { num_expected -= 1; - assert_eq!(val, authorization_data.state); + assert_eq!(val, authorization_data.state.secret().as_str()); } - "nonce" => { + "code_challenge" => { + code_challenge = Some(val); + num_expected -= 1; + } + "code_challenge_method" => { + assert_eq!(val, "S256"); num_expected -= 1; - nonce = Some(val); } _ => panic!("unexpected query parameter: {key}={val}"), } @@ -195,8 +212,11 @@ async fn test_login() -> anyhow::Result<()> { let authorization_data_guard = data.authorization_data.lock().await; let state = authorization_data_guard.get(&authorization_data.state).context("missing state")?; - let nonce = nonce.context("missing nonce")?; - assert_eq!(nonce, state.nonce); + let code_challenge = code_challenge.context("missing code_challenge")?; + assert_eq!( + code_challenge, + PkceCodeChallenge::from_code_verifier_sha256(&state.pkce_verifier).as_str() + ); assert!(authorization_data.url.as_str().starts_with(&issuer)); assert_eq!(authorization_data.url.path(), "/oauth2/authorize"); @@ -217,17 +237,17 @@ fn test_authorization_response() -> anyhow::Result<()> { AuthorizationResponse::parse_uri(&uri), Ok(AuthorizationResponse::Success(AuthorizationCode { code, state })) => { assert_eq!(code, "123"); - assert_eq!(state, "456"); + assert_eq!(state.secret(), "456"); } ); - let uri = Url::parse("https://example.com?error=invalid_grant&state=456")?; + let uri = Url::parse("https://example.com?error=invalid_scope&state=456")?; assert_matches!( AuthorizationResponse::parse_uri(&uri), Ok(AuthorizationResponse::Error(AuthorizationError { error, state })) => { - assert_eq!(error.error, ClientErrorCode::InvalidGrant); - assert_eq!(error.error_description, None); - assert_eq!(state, "456"); + assert_eq!(*error.error(), AuthorizationCodeErrorResponseType::InvalidScope); + assert_eq!(error.error_description(), None); + assert_eq!(state.secret(), "456"); } ); @@ -247,27 +267,30 @@ async fn test_finish_authorization() -> anyhow::Result<()> { // If the state is missing, then any attempt to finish authorizing will fail. let res = oidc - .finish_authorization(AuthorizationCode { code: "42".to_owned(), state: "none".to_owned() }) + .finish_authorization(AuthorizationCode { + code: "42".to_owned(), + state: CsrfToken::new("none".to_owned()), + }) .await; - assert_matches!(res, Err(OidcError::InvalidState)); + assert_matches!( + res, + Err(OidcError::AuthorizationCode(OauthAuthorizationCodeError::InvalidState)) + ); assert!(oidc.session_tokens().is_none()); // Assuming a non-empty state "123"... - let state = "state".to_owned(); + let state = CsrfToken::new("state".to_owned()); let redirect_uri = REDIRECT_URI_STRING; + let (_pkce_code_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256(); let auth_validation_data = AuthorizationValidationData { - state: state.clone(), - nonce: "nonce".to_owned(), - redirect_uri: Url::parse(redirect_uri)?, - code_challenge_verifier: None, + redirect_uri: RedirectUrl::new(redirect_uri.to_owned())?, + pkce_verifier, }; { let data = oidc.data().context("missing data")?; - let prev = data.authorization_data.lock().await.insert(state.clone(), { - AuthorizationValidationData { ..auth_validation_data.clone() } - }); + let prev = data.authorization_data.lock().await.insert(state.clone(), auth_validation_data); assert!(prev.is_none()); } @@ -275,11 +298,14 @@ async fn test_finish_authorization() -> anyhow::Result<()> { let res = oidc .finish_authorization(AuthorizationCode { code: "1337".to_owned(), - state: "none".to_owned(), + state: CsrfToken::new("none".to_owned()), }) .await; - assert_matches!(res, Err(OidcError::InvalidState)); + assert_matches!( + res, + Err(OidcError::AuthorizationCode(OauthAuthorizationCodeError::InvalidState)) + ); assert!(oidc.session_tokens().is_none()); assert!(oidc.data().unwrap().authorization_data.lock().await.get(&state).is_some()); diff --git a/examples/oidc_cli/src/main.rs b/examples/oidc_cli/src/main.rs index 555b573c9..52ccd9162 100644 --- a/examples/oidc_cli/src/main.rs +++ b/examples/oidc_cli/src/main.rs @@ -38,7 +38,8 @@ use matrix_sdk::{ registration::{ClientMetadata, Localized, VerifiedClientMetadata}, requests::GrantType, }, - AuthorizationCode, AuthorizationResponse, OidcAuthorizationData, OidcSession, UserSession, + AuthorizationCode, AuthorizationResponse, CsrfToken, OidcAuthorizationData, OidcSession, + UserSession, }, config::SyncSettings, encryption::{recovery::RecoveryState, CrossSigningResetAuthType}, @@ -746,7 +747,7 @@ fn client_metadata() -> VerifiedClientMetadata { /// Returns the code to obtain the access token. async fn use_auth_url( url: &Url, - state: &str, + state: &CsrfToken, data_rx: oneshot::Receiver, signal_tx: oneshot::Sender<()>, ) -> anyhow::Result { @@ -759,8 +760,7 @@ async fn use_auth_url( let code = match AuthorizationResponse::parse_query(&response_query)? { AuthorizationResponse::Success(code) => code, AuthorizationResponse::Error(err) => { - let err = err.error; - return Err(anyhow!("{}: {:?}", err.error, err.error_description)); + return Err(anyhow!(err.error)); } }; @@ -768,7 +768,7 @@ async fn use_auth_url( // wrong, it is an error. Some clients might want to allow several // authorizations at once, in which case the state string can be used to // identify the session that was authorized. - if code.state != state { + if code.state != *state { bail!("State strings don't match") }