diff --git a/crates/handlers/src/oauth2/token.rs b/crates/handlers/src/oauth2/token.rs index 5608ee688..d921fbf7c 100644 --- a/crates/handlers/src/oauth2/token.rs +++ b/crates/handlers/src/oauth2/token.rs @@ -28,12 +28,15 @@ use mas_oidc_client::types::scope::ScopeToken; use mas_policy::Policy; use mas_router::UrlBuilder; use mas_storage::{ - BoxRepository, RepositoryAccess, + BoxRepository, Pagination, RepositoryAccess, oauth2::{ OAuth2AccessTokenRepository, OAuth2AuthorizationGrantRepository, OAuth2RefreshTokenRepository, OAuth2SessionRepository, }, - user::BrowserSessionRepository, + upstream_oauth2::{ + UpstreamOAuthLinkFilter, UpstreamOAuthLinkTokenRepository, UpstreamOAuthProviderRepository, + }, + user::{BrowserSessionRepository, UserRepository}, }; use mas_templates::{DeviceNameContext, TemplateContext, Templates}; use oauth2_types::{ @@ -41,7 +44,7 @@ use oauth2_types::{ pkce::CodeChallengeError, requests::{ AccessTokenRequest, AccessTokenResponse, AuthorizationCodeGrant, ClientCredentialsGrant, - DeviceCodeGrant, GrantType, RefreshTokenGrant, + DeviceCodeGrant, GrantType, RefreshTokenGrant, TokenExchangeGrant, }, scope, }; @@ -157,6 +160,22 @@ pub(crate) enum RouteError { #[error("failed to provision device")] ProvisionDeviceFailed(#[source] anyhow::Error), + + #[error("subject token is invalid or expired")] + SubjectTokenInvalid, + + #[error("upstream provider not found")] + UpstreamProviderNotFound, + + #[error("user has no link to the requested upstream provider")] + NoUpstreamLink, + + #[error("no stored token for this upstream link")] + NoUpstreamToken, + + #[error("failed to refresh upstream token")] + #[expect(dead_code, reason = "constructed once auto-refresh is implemented")] + UpstreamTokenRefreshFailed(#[source] anyhow::Error), } impl IntoResponse for RouteError { @@ -170,6 +189,7 @@ impl IntoResponse for RouteError { | Self::ProvisionDeviceFailed(_) | Self::NoSuchNextRefreshToken { .. } | Self::NoSuchNextAccessToken { .. } + | Self::UpstreamTokenRefreshFailed(_) ); TOKEN_REQUEST_COUNTER.add(1, &[KeyValue::new(RESULT, "error")]); @@ -181,7 +201,8 @@ impl IntoResponse for RouteError { | Self::NoSuchOAuthSession(_) | Self::ProvisionDeviceFailed(_) | Self::NoSuchNextRefreshToken { .. } - | Self::NoSuchNextAccessToken { .. } => ( + | Self::NoSuchNextAccessToken { .. } + | Self::UpstreamTokenRefreshFailed(_) => ( StatusCode::INTERNAL_SERVER_ERROR, Json(ClientError::from(ClientErrorCode::ServerError)), ), @@ -255,6 +276,40 @@ impl IntoResponse for RouteError { StatusCode::BAD_REQUEST, Json(ClientError::from(ClientErrorCode::UnsupportedGrantType)), ), + + Self::SubjectTokenInvalid => ( + StatusCode::BAD_REQUEST, + Json( + ClientError::from(ClientErrorCode::InvalidGrant) + .with_description("subject_token is invalid or expired".to_owned()), + ), + ), + + Self::UpstreamProviderNotFound => ( + StatusCode::BAD_REQUEST, + Json( + ClientError::from(ClientErrorCode::InvalidGrant).with_description( + "no upstream provider found matching the audience".to_owned(), + ), + ), + ), + + Self::NoUpstreamLink => ( + StatusCode::BAD_REQUEST, + Json( + ClientError::from(ClientErrorCode::InvalidGrant).with_description( + "user has no link to the requested upstream provider".to_owned(), + ), + ), + ), + + Self::NoUpstreamToken => ( + StatusCode::BAD_REQUEST, + Json( + ClientError::from(ClientErrorCode::InvalidGrant) + .with_description("no stored token for this upstream link".to_owned()), + ), + ), }; (sentry_event_id, response).into_response() @@ -391,6 +446,19 @@ pub(crate) async fn post( ) .await? } + AccessTokenRequest::TokenExchange(grant) => { + token_exchange_grant( + &clock, + &activity_tracker, + &grant, + &client, + &encrypter, + repo, + policy, + user_agent, + ) + .await? + } _ => { return Err(RouteError::UnsupportedGrantType); } @@ -1043,6 +1111,173 @@ async fn device_code_grant( Ok((params, repo)) } +/// The expected `subject_token_type` for RFC 8693 token exchange. +const ACCESS_TOKEN_TYPE_URN: &str = "urn:ietf:params:oauth:token-type:access_token"; + +#[tracing::instrument( + name = "handlers.oauth2.token.token_exchange", + fields(client.id = %client.id), + skip_all, +)] +async fn token_exchange_grant( + clock: &impl Clock, + activity_tracker: &BoundActivityTracker, + grant: &TokenExchangeGrant, + client: &Client, + encrypter: &Encrypter, + mut repo: BoxRepository, + mut policy: Policy, + user_agent: Option, +) -> Result<(AccessTokenResponse, BoxRepository), RouteError> { + // The client must be registered for the token exchange grant type + if !client.grant_types.contains(&GrantType::TokenExchange) { + return Err(RouteError::UnauthorizedClient(client.id)); + } + + // 1. Validate subject_token_type + if grant.subject_token_type != ACCESS_TOKEN_TYPE_URN { + return Err(RouteError::BadRequest); + } + + // 2. Parse and validate the subject_token as a MAS access token + let token_type = + TokenType::check(&grant.subject_token).map_err(|_| RouteError::SubjectTokenInvalid)?; + + if token_type != TokenType::AccessToken { + return Err(RouteError::SubjectTokenInvalid); + } + + // 3. Look up the access token and its session + let access_token = repo + .oauth2_access_token() + .find_by_token(&grant.subject_token) + .await? + .ok_or(RouteError::SubjectTokenInvalid)?; + + if !access_token.is_valid(clock.now()) { + return Err(RouteError::SubjectTokenInvalid); + } + + let session = repo + .oauth2_session() + .lookup(access_token.session_id) + .await? + .ok_or(RouteError::NoSuchOAuthSession(access_token.session_id))?; + + if !session.is_valid() { + return Err(RouteError::SubjectTokenInvalid); + } + + // Token exchange requires a user + let user_id = session.user_id.ok_or(RouteError::SubjectTokenInvalid)?; + let user = repo + .user() + .lookup(user_id) + .await? + .ok_or(RouteError::SubjectTokenInvalid)?; + + if !user.is_valid() { + return Err(RouteError::SubjectTokenInvalid); + } + + // 4. Resolve the upstream provider from audience + let audience = grant + .audience + .as_deref() + .ok_or(RouteError::UpstreamProviderNotFound)?; + + let provider = if let Ok(id) = audience.parse::() { + // Try as a ULID + repo.upstream_oauth_provider() + .lookup(id) + .await? + .filter(mas_data_model::UpstreamOAuthProvider::enabled) + } else { + // Try as an issuer URL + repo.upstream_oauth_provider() + .find_by_issuer(audience) + .await? + } + .ok_or(RouteError::UpstreamProviderNotFound)?; + + // 5. Evaluate policy + let scope = grant.scope.clone().unwrap_or_else(|| "".parse().unwrap()); + + let provider_id_str = provider.id.to_string(); + let res = policy + .evaluate_authorization_grant(mas_policy::AuthorizationGrantInput { + user: Some(&user), + client, + session_counts: None, + scope: &scope, + grant_type: mas_policy::GrantType::TokenExchange, + upstream_provider: Some(mas_policy::UpstreamProviderInfo { + id: &provider_id_str, + issuer: provider.issuer.as_deref(), + human_name: provider.human_name.as_deref(), + }), + requester: mas_policy::Requester { + ip_address: activity_tracker.ip(), + user_agent, + }, + }) + .await?; + + if !res.valid() { + return Err(RouteError::DeniedByPolicy(res)); + } + + // 6. Find the upstream link for this user + provider + let filter = UpstreamOAuthLinkFilter::new() + .for_user(&user) + .for_provider(&provider); + + let page = repo + .upstream_oauth_link() + .list(filter, Pagination::first(1)) + .await?; + + let link = page + .edges + .into_iter() + .next() + .map(|edge| edge.node) + .ok_or(RouteError::NoUpstreamLink)?; + + // 7. Find the stored token for this link + let link_token = repo + .upstream_oauth_link_token() + .find_by_link(&link) + .await? + .ok_or(RouteError::NoUpstreamToken)?; + + // 8. Decrypt the upstream access token + let decrypted = encrypter + .decrypt_string(&link_token.encrypted_access_token) + .map_err(|e| RouteError::Internal(Box::new(e)))?; + let upstream_access_token = + String::from_utf8(decrypted).map_err(|e| RouteError::Internal(Box::new(e)))?; + + // 9. Build the response + let mut response = AccessTokenResponse::new(upstream_access_token); + response.issued_token_type = Some(ACCESS_TOKEN_TYPE_URN.to_owned()); + + if let Some(expires_at) = link_token.access_token_expires_at { + let remaining = expires_at - clock.now(); + if remaining > Duration::zero() { + response = response.with_expires_in(remaining); + } + } + + if let Some(ref token_scope) = link_token.token_scope + && let Ok(parsed_scope) = token_scope.parse() + { + response = response.with_scope(parsed_scope); + } + + Ok((response, repo)) +} + #[cfg(test)] mod tests { use hyper::Request; diff --git a/crates/storage-pg/.sqlx/query-b917495985e3d3f2ad270fe309186108fc7ecb310e2f8110c5403346a1ee970e.json b/crates/storage-pg/.sqlx/query-b917495985e3d3f2ad270fe309186108fc7ecb310e2f8110c5403346a1ee970e.json new file mode 100644 index 000000000..b91b65480 --- /dev/null +++ b/crates/storage-pg/.sqlx/query-b917495985e3d3f2ad270fe309186108fc7ecb310e2f8110c5403346a1ee970e.json @@ -0,0 +1,172 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT\n upstream_oauth_provider_id,\n issuer,\n human_name,\n brand_name,\n scope,\n client_id,\n encrypted_client_secret,\n token_endpoint_signing_alg,\n token_endpoint_auth_method,\n id_token_signed_response_alg,\n fetch_userinfo,\n userinfo_signed_response_alg,\n created_at,\n disabled_at,\n claims_imports as \"claims_imports: Json\",\n jwks_uri_override,\n authorization_endpoint_override,\n token_endpoint_override,\n userinfo_endpoint_override,\n discovery_mode,\n pkce_mode,\n response_mode,\n additional_parameters as \"additional_parameters: Json>\",\n forward_login_hint,\n on_backchannel_logout,\n registration_token_required\n FROM upstream_oauth_providers\n WHERE issuer = $1\n AND disabled_at IS NULL\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "upstream_oauth_provider_id", + "type_info": "Uuid" + }, + { + "ordinal": 1, + "name": "issuer", + "type_info": "Text" + }, + { + "ordinal": 2, + "name": "human_name", + "type_info": "Text" + }, + { + "ordinal": 3, + "name": "brand_name", + "type_info": "Text" + }, + { + "ordinal": 4, + "name": "scope", + "type_info": "Text" + }, + { + "ordinal": 5, + "name": "client_id", + "type_info": "Text" + }, + { + "ordinal": 6, + "name": "encrypted_client_secret", + "type_info": "Text" + }, + { + "ordinal": 7, + "name": "token_endpoint_signing_alg", + "type_info": "Text" + }, + { + "ordinal": 8, + "name": "token_endpoint_auth_method", + "type_info": "Text" + }, + { + "ordinal": 9, + "name": "id_token_signed_response_alg", + "type_info": "Text" + }, + { + "ordinal": 10, + "name": "fetch_userinfo", + "type_info": "Bool" + }, + { + "ordinal": 11, + "name": "userinfo_signed_response_alg", + "type_info": "Text" + }, + { + "ordinal": 12, + "name": "created_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 13, + "name": "disabled_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 14, + "name": "claims_imports: Json", + "type_info": "Jsonb" + }, + { + "ordinal": 15, + "name": "jwks_uri_override", + "type_info": "Text" + }, + { + "ordinal": 16, + "name": "authorization_endpoint_override", + "type_info": "Text" + }, + { + "ordinal": 17, + "name": "token_endpoint_override", + "type_info": "Text" + }, + { + "ordinal": 18, + "name": "userinfo_endpoint_override", + "type_info": "Text" + }, + { + "ordinal": 19, + "name": "discovery_mode", + "type_info": "Text" + }, + { + "ordinal": 20, + "name": "pkce_mode", + "type_info": "Text" + }, + { + "ordinal": 21, + "name": "response_mode", + "type_info": "Text" + }, + { + "ordinal": 22, + "name": "additional_parameters: Json>", + "type_info": "Jsonb" + }, + { + "ordinal": 23, + "name": "forward_login_hint", + "type_info": "Bool" + }, + { + "ordinal": 24, + "name": "on_backchannel_logout", + "type_info": "Text" + }, + { + "ordinal": 25, + "name": "registration_token_required", + "type_info": "Bool" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + true, + true, + true, + false, + false, + true, + true, + false, + false, + false, + true, + false, + true, + false, + true, + true, + true, + true, + false, + false, + true, + true, + false, + false, + false + ] + }, + "hash": "b917495985e3d3f2ad270fe309186108fc7ecb310e2f8110c5403346a1ee970e" +} diff --git a/crates/storage-pg/src/upstream_oauth2/provider.rs b/crates/storage-pg/src/upstream_oauth2/provider.rs index 291ba2ff7..e9889036b 100644 --- a/crates/storage-pg/src/upstream_oauth2/provider.rs +++ b/crates/storage-pg/src/upstream_oauth2/provider.rs @@ -1,4 +1,4 @@ -// Copyright 2026 Element Creations Ltd. +// Copyright 2025, 2026 Element Creations Ltd. // Copyright 2024, 2025 New Vector Ltd. // Copyright 2022-2024 The Matrix.org Foundation C.I.C. // @@ -1003,4 +1003,65 @@ impl UpstreamOAuthProviderRepository for PgUpstreamOAuthProviderRepository<'_> { let res: Result, _> = res.into_iter().map(TryInto::try_into).collect(); Ok(res?) } + + #[tracing::instrument( + name = "db.upstream_oauth_provider.find_by_issuer", + skip_all, + fields( + db.query.text, + upstream_oauth_provider.issuer = issuer, + ), + err, + )] + async fn find_by_issuer( + &mut self, + issuer: &str, + ) -> Result, Self::Error> { + let res = sqlx::query_as!( + ProviderLookup, + r#" + SELECT + upstream_oauth_provider_id, + issuer, + human_name, + brand_name, + scope, + client_id, + encrypted_client_secret, + token_endpoint_signing_alg, + token_endpoint_auth_method, + id_token_signed_response_alg, + fetch_userinfo, + userinfo_signed_response_alg, + created_at, + disabled_at, + claims_imports as "claims_imports: Json", + jwks_uri_override, + authorization_endpoint_override, + token_endpoint_override, + userinfo_endpoint_override, + discovery_mode, + pkce_mode, + response_mode, + additional_parameters as "additional_parameters: Json>", + forward_login_hint, + on_backchannel_logout, + registration_token_required + FROM upstream_oauth_providers + WHERE issuer = $1 + AND disabled_at IS NULL + "#, + issuer, + ) + .traced() + .fetch_optional(&mut *self.conn) + .await?; + + let res = res + .map(UpstreamOAuthProvider::try_from) + .transpose() + .map_err(DatabaseError::from)?; + + Ok(res) + } } diff --git a/crates/storage/src/upstream_oauth2/provider.rs b/crates/storage/src/upstream_oauth2/provider.rs index df217b433..011f5b101 100644 --- a/crates/storage/src/upstream_oauth2/provider.rs +++ b/crates/storage/src/upstream_oauth2/provider.rs @@ -1,3 +1,4 @@ +// Copyright 2025, 2026 Element Creations Ltd. // Copyright 2024, 2025 New Vector Ltd. // Copyright 2022-2024 The Matrix.org Foundation C.I.C. // @@ -287,6 +288,22 @@ pub trait UpstreamOAuthProviderRepository: Send + Sync { /// /// Returns [`Self::Error`] if the underlying repository fails async fn all_enabled(&mut self) -> Result, Self::Error>; + + /// Lookup an enabled upstream OAuth provider by its issuer URL + /// + /// Returns `None` if no enabled provider with the given issuer was found + /// + /// # Parameters + /// + /// * `issuer`: The issuer URL to look up + /// + /// # Errors + /// + /// Returns [`Self::Error`] if the underlying repository fails + async fn find_by_issuer( + &mut self, + issuer: &str, + ) -> Result, Self::Error>; } repository_impl!(UpstreamOAuthProviderRepository: @@ -328,4 +345,9 @@ repository_impl!(UpstreamOAuthProviderRepository: ) -> Result; async fn all_enabled(&mut self) -> Result, Self::Error>; + + async fn find_by_issuer( + &mut self, + issuer: &str, + ) -> Result, Self::Error>; );