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
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,10 @@ from google.auth.exceptions import MutualTLSChannelError
from google.protobuf import json_format
from urllib.parse import urlparse, urlunparse

from google.auth.transport import mtls # type: ignore

{# TODO: Remove client cert compatibility fallbacks when the minimum supported
version of google-auth is >= 2.43.0 (currently 2.14.1+). #}
try:
# note: `#type: ignore` is added because the return type for `should_use_client_cert`
# is different than that of the fallback implementation below. This will be removed once
Expand All @@ -50,6 +54,26 @@ except ImportError: # pragma: NO COVER
)
return use_client_cert == "true"


def get_client_cert_source(provided_cert_source, use_cert_flag):
"""Return the client cert source to be used by the client.

Args:
provided_cert_source (bytes): The client certificate source provided.
use_cert_flag (bool): A flag indicating whether to use the client certificate.

Returns:
bytes or None: The client cert source to be used by the client.
"""
client_cert_source = None
if use_cert_flag:
if provided_cert_source:
client_cert_source = provided_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
return client_cert_source
Comment thread
hebaalazzeh marked this conversation as resolved.


DEFAULT_UNIVERSE = "googleapis.com"


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,13 +30,12 @@ from google.api_core import exceptions as core_exceptions
from google.api_core import extended_operation
{% endif %}
from google.api_core import gapic_v1
from {{package_path}}._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert
from {{package_path}}._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, get_client_cert_source
{% if has_auto_populated_fields %}
from {{package_path}}._compat import setup_request_id
{% endif %}
from google.api_core import retry as retries
from google.auth import credentials as ga_credentials # type: ignore
from google.auth.transport import mtls # type: ignore
from google.auth.transport.grpc import SslCredentials # type: ignore
from google.auth.exceptions import MutualTLSChannelError # type: ignore
from google.oauth2 import service_account # type: ignore
Expand Down Expand Up @@ -276,12 +275,9 @@ class {{ service.client_name }}(metaclass={{ service.client_name }}Meta):
raise MutualTLSChannelError("Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`")

# Figure out the client cert source to use.
client_cert_source = None
if use_client_cert:
if client_options.client_cert_source:
client_cert_source = client_options.client_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
client_cert_source = get_client_cert_source(
client_options.client_cert_source, use_client_cert
)

# Figure out which api endpoint to use.
if client_options.api_endpoint is not None:
Expand Down Expand Up @@ -314,24 +310,7 @@ class {{ service.client_name }}(metaclass={{ service.client_name }}Meta):
raise MutualTLSChannelError("Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`")
return use_client_cert, use_mtls_endpoint, universe_domain_env

@staticmethod
def _get_client_cert_source(provided_cert_source, use_cert_flag):
"""Return the client cert source to be used by the client.

Args:
provided_cert_source (bytes): The client certificate source provided.
use_cert_flag (bool): A flag indicating whether to use the client certificate.

Returns:
bytes or None: The client cert source to be used by the client.
"""
client_cert_source = None
if use_cert_flag:
if provided_cert_source:
client_cert_source = provided_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
return client_cert_source


def _validate_universe_domain(self):
Expand Down Expand Up @@ -460,7 +439,7 @@ class {{ service.client_name }}(metaclass={{ service.client_name }}Meta):
universe_domain_opt = getattr(self._client_options, 'universe_domain', None)

self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = {{ service.client_name }}._read_environment_variables()
self._client_cert_source = {{ service.client_name }}._get_client_cert_source(self._client_options.client_cert_source, self._use_client_cert)
self._client_cert_source = get_client_cert_source(self._client_options.client_cert_source, self._use_client_cert)
self._universe_domain = get_universe_domain(universe_domain_opt, self._universe_domain_env)
self._api_endpoint: str = "" # updated below, depending on `transport`

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -168,18 +168,6 @@ def set_event_loop():
asyncio.set_event_loop(None)


def test__get_client_cert_source():
mock_provided_cert_source = mock.Mock()
mock_default_cert_source = mock.Mock()

assert {{ service.client_name }}._get_client_cert_source(None, False) is None
assert {{ service.client_name }}._get_client_cert_source(mock_provided_cert_source, False) is None
assert {{ service.client_name }}._get_client_cert_source(mock_provided_cert_source, True) == mock_provided_cert_source

with mock.patch('google.auth.transport.mtls.has_default_client_cert_source', return_value=True):
with mock.patch('google.auth.transport.mtls.default_client_cert_source', return_value=mock_default_cert_source):
assert {{ service.client_name }}._get_client_cert_source(None, True) is mock_default_cert_source
assert {{ service.client_name }}._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source



Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ import google.auth.transport.mtls

{% set package_path = api.naming.module_namespace|join('.') + "." + api.naming.versioned_module_name %}
from {{package_path}}._compat import transcode_request
from {{package_path}}._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert
from {{package_path}}._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, get_client_cert_source
{% if has_auto_populated_fields %}
from {{package_path}}._compat import setup_request_id
{% endif %}
Expand Down Expand Up @@ -497,4 +497,18 @@ def test_transcode_request_proto_plus_wrapper():
transcoded, _, _ = transcode_request(http_options, mock_proto_plus)
assert transcoded["uri"] == "/v1/test/proto-plus-field"


def test_get_client_cert_source():
mock_provided_cert_source = mock.Mock()
mock_default_cert_source = mock.Mock()

assert get_client_cert_source(None, False) is None
assert get_client_cert_source(mock_provided_cert_source, False) is None
assert get_client_cert_source(mock_provided_cert_source, True) == mock_provided_cert_source

with mock.patch("google.auth.transport.mtls.has_default_client_cert_source", return_value=True):
with mock.patch("google.auth.transport.mtls.default_client_cert_source", return_value=mock_default_cert_source):
assert get_client_cert_source(None, True) is mock_default_cert_source
assert get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

If get_client_cert_source is type-hinted with use_cert_flag: bool, passing the string "true" here can cause static type checkers (like mypy) to fail when type-checking the generated test files. It is safer and more idiomatic to use the boolean True instead.

            assert get_client_cert_source(mock_provided_cert_source, True) is mock_provided_cert_source


{% endblock %}
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,8 @@
from google.protobuf import json_format
from urllib.parse import urlparse, urlunparse

from google.auth.transport import mtls # type: ignore

try:
# note: `#type: ignore` is added because the return type for `should_use_client_cert`
# is different than that of the fallback implementation below. This will be removed once
Expand All @@ -42,6 +44,26 @@ def should_use_client_cert():
)
return use_client_cert == "true"


def get_client_cert_source(provided_cert_source, use_cert_flag):
"""Return the client cert source to be used by the client.

Args:
provided_cert_source (bytes): The client certificate source provided.
use_cert_flag (bool): A flag indicating whether to use the client certificate.

Returns:
bytes or None: The client cert source to be used by the client.
"""
client_cert_source = None
if use_cert_flag:
if provided_cert_source:
client_cert_source = provided_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
return client_cert_source


DEFAULT_UNIVERSE = "googleapis.com"


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,10 +27,9 @@
from google.api_core import client_options as client_options_lib
from google.api_core import exceptions as core_exceptions
from google.api_core import gapic_v1
from google.cloud.asset_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert
from google.cloud.asset_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, get_client_cert_source
from google.api_core import retry as retries
from google.auth import credentials as ga_credentials # type: ignore
from google.auth.transport import mtls # type: ignore
from google.auth.transport.grpc import SslCredentials # type: ignore
from google.auth.exceptions import MutualTLSChannelError # type: ignore
from google.oauth2 import service_account # type: ignore
Expand Down Expand Up @@ -331,12 +330,9 @@ def get_mtls_endpoint_and_cert_source(cls, client_options: Optional[client_optio
raise MutualTLSChannelError("Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`")

# Figure out the client cert source to use.
client_cert_source = None
if use_client_cert:
if client_options.client_cert_source:
client_cert_source = client_options.client_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
client_cert_source = get_client_cert_source(
client_options.client_cert_source, use_client_cert
)

# Figure out which api endpoint to use.
if client_options.api_endpoint is not None:
Expand Down Expand Up @@ -369,25 +365,6 @@ def _read_environment_variables():
raise MutualTLSChannelError("Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`")
return use_client_cert, use_mtls_endpoint, universe_domain_env

@staticmethod
def _get_client_cert_source(provided_cert_source, use_cert_flag):
"""Return the client cert source to be used by the client.

Args:
provided_cert_source (bytes): The client certificate source provided.
use_cert_flag (bool): A flag indicating whether to use the client certificate.

Returns:
bytes or None: The client cert source to be used by the client.
"""
client_cert_source = None
if use_cert_flag:
if provided_cert_source:
client_cert_source = provided_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
return client_cert_source

def _validate_universe_domain(self):
"""Validates client's and credentials' universe domains are consistent.

Expand Down Expand Up @@ -511,7 +488,7 @@ def __init__(self, *,
universe_domain_opt = getattr(self._client_options, 'universe_domain', None)

self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = AssetServiceClient._read_environment_variables()
self._client_cert_source = AssetServiceClient._get_client_cert_source(self._client_options.client_cert_source, self._use_client_cert)
self._client_cert_source = get_client_cert_source(self._client_options.client_cert_source, self._use_client_cert)
self._universe_domain = get_universe_domain(universe_domain_opt, self._universe_domain_env)
self._api_endpoint: str = "" # updated below, depending on `transport`

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -132,20 +132,6 @@ def set_event_loop():
asyncio.set_event_loop(None)


def test__get_client_cert_source():
mock_provided_cert_source = mock.Mock()
mock_default_cert_source = mock.Mock()

assert AssetServiceClient._get_client_cert_source(None, False) is None
assert AssetServiceClient._get_client_cert_source(mock_provided_cert_source, False) is None
assert AssetServiceClient._get_client_cert_source(mock_provided_cert_source, True) == mock_provided_cert_source

with mock.patch('google.auth.transport.mtls.has_default_client_cert_source', return_value=True):
with mock.patch('google.auth.transport.mtls.default_client_cert_source', return_value=mock_default_cert_source):
assert AssetServiceClient._get_client_cert_source(None, True) is mock_default_cert_source
assert AssetServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source


@pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [
(401, CRED_INFO_JSON, True),
(403, CRED_INFO_JSON, True),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
import google.auth.transport.mtls

from google.cloud.asset_v1._compat import transcode_request
from google.cloud.asset_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert
from google.cloud.asset_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, get_client_cert_source

from google.auth.exceptions import MutualTLSChannelError
from google.api_core.universe import EmptyUniverseError
Expand Down Expand Up @@ -411,3 +411,17 @@ def test_transcode_request_proto_plus_wrapper():

transcoded, _, _ = transcode_request(http_options, mock_proto_plus)
assert transcoded["uri"] == "/v1/test/proto-plus-field"


def test_get_client_cert_source():
mock_provided_cert_source = mock.Mock()
mock_default_cert_source = mock.Mock()

assert get_client_cert_source(None, False) is None
assert get_client_cert_source(mock_provided_cert_source, False) is None
assert get_client_cert_source(mock_provided_cert_source, True) == mock_provided_cert_source

with mock.patch("google.auth.transport.mtls.has_default_client_cert_source", return_value=True):
with mock.patch("google.auth.transport.mtls.default_client_cert_source", return_value=mock_default_cert_source):
assert get_client_cert_source(None, True) is mock_default_cert_source
assert get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,8 @@
from google.protobuf import json_format
from urllib.parse import urlparse, urlunparse

from google.auth.transport import mtls # type: ignore

try:
# note: `#type: ignore` is added because the return type for `should_use_client_cert`
# is different than that of the fallback implementation below. This will be removed once
Expand All @@ -42,6 +44,26 @@ def should_use_client_cert():
)
return use_client_cert == "true"


def get_client_cert_source(provided_cert_source, use_cert_flag):
"""Return the client cert source to be used by the client.

Args:
provided_cert_source (bytes): The client certificate source provided.
use_cert_flag (bool): A flag indicating whether to use the client certificate.

Returns:
bytes or None: The client cert source to be used by the client.
"""
client_cert_source = None
if use_cert_flag:
if provided_cert_source:
client_cert_source = provided_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
return client_cert_source


DEFAULT_UNIVERSE = "googleapis.com"


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,10 +27,9 @@
from google.api_core import client_options as client_options_lib
from google.api_core import exceptions as core_exceptions
from google.api_core import gapic_v1
from google.iam.credentials_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert
from google.iam.credentials_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, get_client_cert_source
from google.api_core import retry as retries
from google.auth import credentials as ga_credentials # type: ignore
from google.auth.transport import mtls # type: ignore
from google.auth.transport.grpc import SslCredentials # type: ignore
from google.auth.exceptions import MutualTLSChannelError # type: ignore
from google.oauth2 import service_account # type: ignore
Expand Down Expand Up @@ -268,12 +267,9 @@ def get_mtls_endpoint_and_cert_source(cls, client_options: Optional[client_optio
raise MutualTLSChannelError("Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`")

# Figure out the client cert source to use.
client_cert_source = None
if use_client_cert:
if client_options.client_cert_source:
client_cert_source = client_options.client_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
client_cert_source = get_client_cert_source(
client_options.client_cert_source, use_client_cert
)

# Figure out which api endpoint to use.
if client_options.api_endpoint is not None:
Expand Down Expand Up @@ -306,25 +302,6 @@ def _read_environment_variables():
raise MutualTLSChannelError("Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`")
return use_client_cert, use_mtls_endpoint, universe_domain_env

@staticmethod
def _get_client_cert_source(provided_cert_source, use_cert_flag):
"""Return the client cert source to be used by the client.

Args:
provided_cert_source (bytes): The client certificate source provided.
use_cert_flag (bool): A flag indicating whether to use the client certificate.

Returns:
bytes or None: The client cert source to be used by the client.
"""
client_cert_source = None
if use_cert_flag:
if provided_cert_source:
client_cert_source = provided_cert_source
elif mtls.has_default_client_cert_source():
client_cert_source = mtls.default_client_cert_source()
return client_cert_source

def _validate_universe_domain(self):
"""Validates client's and credentials' universe domains are consistent.

Expand Down Expand Up @@ -448,7 +425,7 @@ def __init__(self, *,
universe_domain_opt = getattr(self._client_options, 'universe_domain', None)

self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = IAMCredentialsClient._read_environment_variables()
self._client_cert_source = IAMCredentialsClient._get_client_cert_source(self._client_options.client_cert_source, self._use_client_cert)
self._client_cert_source = get_client_cert_source(self._client_options.client_cert_source, self._use_client_cert)
self._universe_domain = get_universe_domain(universe_domain_opt, self._universe_domain_env)
self._api_endpoint: str = "" # updated below, depending on `transport`

Expand Down
Loading
Loading