Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 12 additions & 1 deletion packages/google-auth/google/auth/transport/_mtls_helper.py
Original file line number Diff line number Diff line change
Expand Up @@ -836,10 +836,21 @@ def check_parameters_for_unauthorized_response(cached_cert):


def call_client_cert_callback():
"""Calls the client cert callback and returns the certificate and key."""
"""Calls the client cert callback and returns the certificate and key.

If the cert provider returns a passphrase-protected private key, it is
decrypted before being returned, so callers always receive an unencrypted
PEM key that can be passed directly to TLS libraries (e.g. gRPC).

Returns:
Tuple[bytes, bytes]: The client certificate and (unencrypted) private
key bytes in PEM format.
"""
_, cert_bytes, key_bytes, passphrase = get_client_ssl_credentials(
generate_encrypted_key=True
)
if passphrase is not None:
key_bytes = decrypt_private_key(key_bytes, passphrase)
return cert_bytes, key_bytes


Expand Down
34 changes: 29 additions & 5 deletions packages/google-auth/google/auth/transport/grpc.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,13 @@

from __future__ import absolute_import

import functools
import logging
import warnings
from typing import Optional

from google.auth import exceptions
from google.auth.transport import _mtls_helper, mtls
from google.auth.transport import _mtls_helper, mtls, mtls_interceptor
from google.oauth2 import service_account

try:
Expand Down Expand Up @@ -282,6 +284,7 @@ def my_client_cert_callback():
)

# If SSL credentials are not explicitly set, try client_cert_callback and ADC.
cached_cert: Optional[bytes] = None
if not ssl_credentials:
use_client_cert = _mtls_helper.check_use_client_cert()
if use_client_cert and client_cert_callback:
Expand All @@ -290,19 +293,38 @@ def my_client_cert_callback():
ssl_credentials = grpc.ssl_channel_credentials(
certificate_chain=cert, private_key=key
)
cached_cert = cert
elif use_client_cert:
# Use application default SSL credentials.
adc_ssl_credentils = SslCredentials()
ssl_credentials = adc_ssl_credentils.ssl_credentials
adc_ssl_credentials = SslCredentials()
ssl_credentials = adc_ssl_credentials.ssl_credentials
cached_cert = adc_ssl_credentials._cached_cert
else:
ssl_credentials = grpc.ssl_channel_credentials()

# Combine the ssl credentials and the authorization credentials.
composite_credentials = grpc.composite_channel_credentials(
ssl_credentials, google_auth_credentials
)

return grpc.secure_channel(target, composite_credentials, **kwargs)
is_recreation = kwargs.pop("_is_recreation", False)
channel = grpc.secure_channel(target, composite_credentials, **kwargs)
# Avoid wrapping if mTLS is disabled or if this is a channel recreation call
if cached_cert and not is_recreation:
# Package arguments so the channel can be recreated later
create_channel_fn = functools.partial(
secure_authorized_channel,
credentials=credentials,
request=request,
target=target,
_is_recreation=True, # Hidden flag to stop recursion
**kwargs,
)
wrapper = mtls_interceptor.MTLSRefreshingChannel(
target, create_channel_fn, channel, cached_cert
)
interceptor = mtls_interceptor.CertRotationInterceptor(wrapper=wrapper)
return grpc.intercept_channel(wrapper, interceptor)
return channel


class SslCredentials:
Expand All @@ -326,6 +348,7 @@ class SslCredentials:

def __init__(self):
use_client_cert = _mtls_helper.check_use_client_cert()
self._cached_cert = None
if not use_client_cert:
self._is_mtls = False
else:
Expand Down Expand Up @@ -354,6 +377,7 @@ def ssl_credentials(self):
self._ssl_credentials = grpc.ssl_channel_credentials(
certificate_chain=cert, private_key=key
)
self._cached_cert = cert
else:
self._ssl_credentials = grpc.ssl_channel_credentials()
self._is_mtls = False
Expand Down
Loading
Loading