diff --git a/README.md b/README.md index b286c22..67da0dd 100644 --- a/README.md +++ b/README.md @@ -71,11 +71,17 @@ development: auth0_client_id: auth0_client_assertion_signing_key: auth0_client_assertion_signing_algorithm: + auth0_client_assertion_signing_key_id: ``` **Note**: you must upload the corresponding public key to your Auth0 tenant, so that Auth0 is able to verify the JWT signature. client_assertion_signing_algorithm is optional and defaults to RS256. +client_assertion_signing_key_id is optional. When set, it is sent as the `kid` header of the +client assertion JWT, which lets Auth0 pick the matching public key when your tenant has more +than one credential registered — for example while rotating keys. Use the key ID that Auth0 +assigned to the uploaded public key. When it is not set, no `kid` header is sent. + ### Create the initializer Create a new Ruby file in `./config/initializers/auth0.rb` to configure the OmniAuth middleware: @@ -115,7 +121,8 @@ Rails.application.config.middleware.use OmniAuth::Builder do scope: 'openid profile' }, client_assertion_signing_key: OpenSSL::PKey::RSA.new(AUTH0_CONFIG[:auth0_client_assertion_signing_key]), - client_assertion_signing_algorithm: AUTH0_CONFIG[:auth0_client_assertion_signing_algorithm] + client_assertion_signing_algorithm: AUTH0_CONFIG[:auth0_client_assertion_signing_algorithm], + client_assertion_signing_key_id: AUTH0_CONFIG[:auth0_client_assertion_signing_key_id] ) end ``` diff --git a/lib/omniauth/auth0/jwt_token.rb b/lib/omniauth/auth0/jwt_token.rb index 62a1e16..55ab3f9 100644 --- a/lib/omniauth/auth0/jwt_token.rb +++ b/lib/omniauth/auth0/jwt_token.rb @@ -8,21 +8,40 @@ module Auth0 # JWTToken class to generate a JWT token for client assertion # as per the OAuth 2.0 Client Credentials Grant specification. class JWTToken - attr_reader :client_id, :domain_url, :client_assertion_signing_key, :client_assertion_signing_algorithm + attr_reader :client_id, :domain_url, :client_assertion_signing_key, :client_assertion_signing_algorithm, + :client_assertion_signing_key_id - def initialize(client_id, domain_url, client_assertion_signing_key, client_assertion_signing_algorithm = nil) + # Create a new client assertion JWT generator. + # @param client_id string - Application Client ID. + # @param domain_url string - Application domain, used to build the audience. + # @param client_assertion_signing_key key - Private key used to sign the assertion. + # @param client_assertion_signing_algorithm string - Signing algorithm, defaults to RS256. + # @param client_assertion_signing_key_id string - Key ID of the signing key (optional). When + # given, it is sent as the "kid" header so Auth0 can pick the matching public key. + def initialize(client_id, domain_url, client_assertion_signing_key, client_assertion_signing_algorithm = nil, + client_assertion_signing_key_id: nil) @client_id = client_id @domain_url = domain_url @client_assertion_signing_key = client_assertion_signing_key @client_assertion_signing_algorithm = client_assertion_signing_algorithm || 'RS256' + @client_assertion_signing_key_id = client_assertion_signing_key_id end def jwt_token - JWT.encode(jwt_payload, client_assertion_signing_key, client_assertion_signing_algorithm) + JWT.encode(jwt_payload, client_assertion_signing_key, client_assertion_signing_algorithm, jwt_headers) end private + # Build the additional JWT header parameters. The "kid" (key ID) header is only included when a + # key ID is configured, so that Auth0 can select the matching public key during key rotation. + # @return hash - The extra headers to merge into the JWT header, empty hash if none. + def jwt_headers + return {} if ['', nil].include?(client_assertion_signing_key_id) + + { kid: client_assertion_signing_key_id } + end + def jwt_payload { iss: client_id, diff --git a/lib/omniauth/strategies/auth0.rb b/lib/omniauth/strategies/auth0.rb index 89c7d97..981d50b 100644 --- a/lib/omniauth/strategies/auth0.rb +++ b/lib/omniauth/strategies/auth0.rb @@ -150,7 +150,8 @@ def jwt_token OmniAuth::Auth0::JWTToken.new(options.client_id, domain_url, options.client_assertion_signing_key, - options.client_assertion_signing_algorithm) + options.client_assertion_signing_algorithm, + client_assertion_signing_key_id: options.client_assertion_signing_key_id) .jwt_token end diff --git a/spec/omniauth/auth0/jwt_token_spec.rb b/spec/omniauth/auth0/jwt_token_spec.rb index ada4d32..88f79fa 100644 --- a/spec/omniauth/auth0/jwt_token_spec.rb +++ b/spec/omniauth/auth0/jwt_token_spec.rb @@ -44,6 +44,43 @@ expect(decoded_token[0]['jti']).to eq(uuid) end + context 'when a client_assertion_signing_key_id is given' do + it 'includes it as the kid header' do + jwt_token = described_class.new(client_id, + domain_url, + client_assertion_signing_key, + 'RS256', + client_assertion_signing_key_id: 'KEY_ID') + .jwt_token + decoded_token = JWT.decode(jwt_token, client_assertion_signing_key, true, { algorithm: 'RS256' }) + + expect(decoded_token[1]['kid']).to eq('KEY_ID') + expect(decoded_token[1]['alg']).to eq('RS256') + expect(decoded_token[0]['iss']).to eq(client_id) + end + end + + context 'when no client_assertion_signing_key_id is given' do + it 'omits the kid header' do + jwt_token = described_class.new(client_id, domain_url, client_assertion_signing_key).jwt_token + decoded_token = JWT.decode(jwt_token, client_assertion_signing_key, true, { algorithm: 'RS256' }) + + expect(decoded_token[1]).not_to have_key('kid') + end + + it 'omits the kid header when the key id is blank' do + jwt_token = described_class.new(client_id, + domain_url, + client_assertion_signing_key, + 'RS256', + client_assertion_signing_key_id: '') + .jwt_token + decoded_token = JWT.decode(jwt_token, client_assertion_signing_key, true, { algorithm: 'RS256' }) + + expect(decoded_token[1]).not_to have_key('kid') + end + end + context 'when using ES256 algorithm' do let(:client_assertion_signing_key) { OpenSSL::PKey::EC.generate('prime256v1') } diff --git a/spec/omniauth/strategies/auth0_spec.rb b/spec/omniauth/strategies/auth0_spec.rb index 60482fe..816258e 100644 --- a/spec/omniauth/strategies/auth0_spec.rb +++ b/spec/omniauth/strategies/auth0_spec.rb @@ -491,6 +491,11 @@ def stub_auth(body) JWT.encode payload, rsa_private_key, 'RS256', kid: valid_jwks_kid end + # Header of the client assertion captured by the last stub_auth request. + def client_assertion_headers + JWT.decode(@client_assertion, nil, false)[1] + end + def jwt_token?(token) JWT.decode(token, nil, false) true @@ -528,6 +533,7 @@ def stub_auth(body, stubbed_jwt_token: true) .with do |request| params = URI.decode_www_form(request.body).to_h token = params['client_assertion'] + @client_assertion = token request.headers['Auth0-Client'] == telemetry_value && params['grant_type'] == described_class::AUTHORIZATION_CODE_GRANT_TYPE && @@ -550,12 +556,13 @@ def stub_expected_jwks ) end - def stub_jwt_token(algorithm: client_assertion_signing_algorithm) + def stub_jwt_token(algorithm: client_assertion_signing_algorithm, key_id: nil) allow(OmniAuth::Auth0::JWTToken).to receive(:new) .with(client_id, domain_url, client_assertion_signing_key, - algorithm) + algorithm, + client_assertion_signing_key_id: key_id) .and_return(instance_double(OmniAuth::Auth0::JWTToken, jwt_token: jwt_token)) end @@ -582,6 +589,41 @@ def stub_jwt_token(algorithm: client_assertion_signing_algorithm) it_behaves_like 'basic oauth callback assertions' end + context 'basic oauth w/client assertion signing key id' do + let(:client_assertion_signing_key_id) { 'CLIENT_ASSERTION_SIGNING_KEY_ID' } + + before do + @app = make_application(client_secret: nil, + client_assertion_signing_key: client_assertion_signing_key, + client_assertion_signing_key_id: client_assertion_signing_key_id) + stub_jwt_token(algorithm: nil, key_id: client_assertion_signing_key_id) + stub_auth(oauth_response) + stub_userinfo(basic_user_info) + trigger_callback + end + + it_behaves_like 'basic oauth callback assertions' + end + + context 'basic oauth w/client assertion signing key id, without stubbing jwt token' do + let(:client_assertion_signing_key_id) { 'CLIENT_ASSERTION_SIGNING_KEY_ID' } + + before do + @app = make_application(client_secret: nil, + client_assertion_signing_key: client_assertion_signing_key, + client_assertion_signing_key_id: client_assertion_signing_key_id) + stub_auth(oauth_response, stubbed_jwt_token: false) + stub_userinfo(basic_user_info) + trigger_callback + end + + it_behaves_like 'basic oauth callback assertions' + + it 'sends the key id as the kid header of the client assertion' do + expect(client_assertion_headers['kid']).to eq(client_assertion_signing_key_id) + end + end + context 'basic oauth w/refresh token' do before do @app = make_application(client_secret: nil,